From c3a23fe499757e651609d8f8845e578062ab7d74 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 7 Oct 2026 00:05:27 +0000 Subject: [PATCH] refactor: expose core private helpers under public names (#44871) * refactor: expose core private symbols with compatibility aliases Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * types: narrow core migration diagnostics Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix: preserve runtime behavior in core symbol migration Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix: preserve core private usage migration behavior Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix: preserve private value rebinding compatibility Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix: restore optional imports and cover public helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test: exempt router property from call coverage Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test: align recursive detector ignore names Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: mateo Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- enterprise/enterprise_hooks/__init__.py | 8 +- .../enterprise_hooks/banned_keywords.py | 3 +- .../enterprise_hooks/blocked_user_list.py | 3 +- .../google_text_moderation.py | 5 +- .../enterprise_hooks/openai_moderation.py | 3 +- .../send_emails/endpoints.py | 35 +- .../integrations/custom_guardrail.py | 2 +- .../proxy/common_utils/check_batch_cost.py | 82 +- .../proxy/enterprise_routes.py | 4 +- .../proxy/hooks/managed_files.py | 61 +- .../proxy/hooks/managed_vector_stores.py | 21 +- .../litellm_enterprise/proxy/proxy_server.py | 2 +- enterprise/litellm_enterprise/proxy/utils.py | 3 +- .../proxy/vector_stores/endpoints.py | 2 +- litellm/__init__.py | 28 +- litellm/_lazy_imports.py | 22 +- litellm/_lazy_imports_registry.py | 43 +- litellm/_logging.py | 54 +- litellm/_redis.py | 6 +- litellm/_redis_credential_provider.py | 7 +- litellm/batches/batch_utils.py | 101 +- litellm/caching/caching.py | 34 +- litellm/caching/caching_handler.py | 58 +- litellm/caching/redis_cache.py | 22 +- .../transformation.py | 12 +- litellm/constants.py | 2 +- litellm/containers/utils.py | 4 +- litellm/cost_calculator.py | 60 +- litellm/files/main.py | 5 +- litellm/files/types.py | 30 +- .../SlackAlerting/slack_alerting.py | 4 +- litellm/integrations/SlackAlerting/utils.py | 17 +- .../arize/arize_phoenix_prompt_manager.py | 8 +- .../bitbucket/bitbucket_prompt_manager.py | 8 +- litellm/integrations/custom_logger.py | 2 +- litellm/integrations/dotprompt/__init__.py | 9 +- .../integrations/dotprompt/prompt_manager.py | 8 +- litellm/integrations/dynamodb.py | 14 +- litellm/integrations/focus/export_engine.py | 95 +- .../gitlab/gitlab_prompt_manager.py | 26 +- litellm/integrations/langfuse/langfuse.py | 2 +- .../mavvrik_focus/mavvrik_focus_logger.py | 12 +- litellm/integrations/opentelemetry.py | 16 +- litellm/integrations/otel/plumbing/metrics.py | 8 +- litellm/integrations/otel/presets/weave.py | 4 +- litellm/integrations/prometheus.py | 48 +- .../prometheus_helpers/__init__.py | 65 +- litellm/integrations/prometheus_services.py | 10 +- litellm/integrations/weave/weave_otel.py | 9 +- litellm/interactions/streaming_iterator.py | 2 +- litellm/litellm_core_utils/core_helpers.py | 33 +- litellm/litellm_core_utils/dd_tracing.py | 5 +- .../exception_mapping_utils.py | 25 +- .../litellm_core_utils/get_litellm_params.py | 18 +- .../get_llm_provider_logic.py | 7 +- .../litellm_core_utils/get_model_cost_map.py | 10 +- .../health_check_helpers.py | 10 +- .../litellm_core_utils/health_check_utils.py | 8 +- litellm/litellm_core_utils/litellm_logging.py | 323 ++-- .../llm_cost_calc/tool_call_cost_tracking.py | 25 +- .../litellm_core_utils/llm_cost_calc/utils.py | 121 +- .../litellm_core_utils/llm_request_utils.py | 24 +- .../convert_dict_to_response.py | 37 +- .../llm_response_utils/response_metadata.py | 2 +- .../logging_callback_manager.py | 52 +- litellm/litellm_core_utils/logging_utils.py | 20 +- .../litellm_core_utils/model_param_helper.py | 19 +- .../prompt_templates/common_utils.py | 32 +- .../prompt_templates/factory.py | 157 +- .../huggingface_template_handler.py | 25 +- .../litellm_core_utils/realtime_streaming.py | 30 +- .../litellm_core_utils/secret_redaction.py | 20 +- .../sensitive_data_masker.py | 12 +- .../litellm_core_utils/streaming_handler.py | 22 +- litellm/litellm_core_utils/token_counter.py | 43 +- litellm/llms/anthropic/chat/handler.py | 4 +- litellm/llms/anthropic/common_utils.py | 12 +- .../pass_through/messages/handler.py | 2 +- .../pass_through/messages/mcp_handler.py | 8 +- .../pass_through/messages/response_cache.py | 2 +- .../messages/streaming_iterator.py | 2 +- .../pass_through/messages/transformation.py | 4 +- .../responses_adapters/transformation.py | 2 +- litellm/llms/azure/common_utils.py | 4 +- .../llms/azure/image_edit/transformation.py | 4 +- litellm/llms/azure_ai/chat/transformation.py | 10 +- .../azure_ai/image_edit/transformation.py | 6 +- .../image_generation/cost_calculator.py | 4 +- .../llms/azure_ai/rerank/transformation.py | 6 +- litellm/llms/base_llm/base_model_iterator.py | 2 +- .../bedrock/chat/converse_transformation.py | 4 +- .../amazon_deepseek_transformation.py | 4 +- litellm/llms/bedrock/common_utils.py | 4 +- ...n_nova_canvas_image_edit_transformation.py | 8 +- litellm/llms/bedrock/realtime/handler.py | 6 +- .../llms/chatgpt/responses/transformation.py | 4 +- .../codestral/completion/transformation.py | 2 +- litellm/llms/custom_httpx/llm_http_handler.py | 16 +- .../llms/databricks/chat/transformation.py | 12 +- litellm/llms/databricks/streaming_utils.py | 4 +- .../llms/gemini/realtime/transformation.py | 4 +- .../responses/transformation.py | 4 +- .../llms/hosted_vllm/chat/transformation.py | 14 +- .../llms/manus/responses/transformation.py | 10 +- litellm/llms/mistral/chat/transformation.py | 4 +- litellm/llms/ollama/chat/transformation.py | 8 +- .../llms/ollama/completion/transformation.py | 4 +- .../llms/openai/chat/gpt_5_transformation.py | 4 +- .../llms/openai/chat/gpt_transformation.py | 12 +- litellm/llms/openai/realtime/handler.py | 4 +- .../llms/openai/responses/transformation.py | 10 +- litellm/llms/openai_like/model_info.py | 4 +- litellm/llms/sagemaker/common_utils.py | 4 +- litellm/llms/vertex_ai/cost_calculator.py | 10 +- .../llms/vertex_ai/gemini/transformation.py | 12 +- .../vertex_and_google_ai_studio_gemini.py | 20 +- .../multimodal_embeddings/transformation.py | 4 +- litellm/llms/vllm/common_utils.py | 4 +- .../volcengine/responses/transformation.py | 4 +- litellm/llms/watsonx/chat/transformation.py | 6 +- litellm/main.py | 182 +-- litellm/passthrough/main.py | 18 +- .../_experimental/mcp_server/mcp_debug.py | 2 +- .../proxy/agent_endpoints/a2a_endpoints.py | 10 +- litellm/proxy/auth/auth_checks.py | 2 +- litellm/proxy/auth/route_checks.py | 2 +- litellm/proxy/common_request_processing.py | 14 +- litellm/proxy/common_utils/callback_utils.py | 18 +- .../common_utils/prompt_cache_pricing.py | 4 +- .../proxy/container_endpoints/ownership.py | 6 +- .../proxy/credential_endpoints/endpoints.py | 8 +- .../guardrails/auto_router_compression.py | 4 +- .../proxy/guardrails/guardrail_endpoints.py | 28 +- .../guardrail_hooks/bedrock_guardrails.py | 4 +- .../unified_guardrail/unified_guardrail.py | 8 +- .../health_endpoints/_health_endpoints.py | 2 +- litellm/proxy/hooks/batch_rate_limiter.py | 20 +- litellm/proxy/hooks/batch_redis_get.py | 2 +- .../proxy/hooks/dynamic_rate_limiter_v3.py | 4 +- .../proxy/hooks/mcp_semantic_filter/hook.py | 6 +- .../proxy/hooks/parallel_request_limiter.py | 6 +- .../hooks/parallel_request_limiter_v3.py | 8 +- .../proxy/hooks/proxy_track_cost_callback.py | 4 +- .../config_override_endpoints.py | 2 +- .../customer_endpoints.py | 4 +- .../key_management_endpoints.py | 2 +- .../organization_endpoints.py | 4 +- .../router_settings_endpoints.py | 2 +- .../openai_files_endpoints/common_utils.py | 6 +- .../llm_passthrough_endpoints.py | 10 +- ...tex_ai_live_passthrough_logging_handler.py | 2 +- .../pass_through_endpoints.py | 4 +- .../streaming_handler.py | 2 +- litellm/proxy/prompts/prompt_endpoints.py | 10 +- litellm/proxy/proxy_cli.py | 4 +- litellm/proxy/proxy_server.py | 36 +- .../proxy/response_api_endpoints/endpoints.py | 4 +- litellm/proxy/route_llm_request.py | 8 +- .../search_tool_management.py | 10 +- litellm/proxy/spend_tracking/savings.py | 4 +- .../spend_management_endpoints.py | 2 +- .../proxy/spend_tracking/vantage_endpoints.py | 8 +- litellm/proxy/utils.py | 20 +- .../proxy/vector_store_endpoints/endpoints.py | 2 +- .../management_endpoints.py | 2 +- litellm/realtime_api/main.py | 13 +- .../session_handler.py | 2 +- .../streaming_iterator.py | 10 +- .../transformation.py | 26 +- litellm/responses/main.py | 72 +- .../responses/mcp/chat_completions_handler.py | 20 +- .../mcp/litellm_proxy_mcp_handler.py | 90 +- .../responses/mcp/mcp_streaming_iterator.py | 16 +- litellm/responses/mcp/request_context.py | 12 +- litellm/responses/sse_output_recovery.py | 2 +- litellm/responses/streaming_iterator.py | 75 +- litellm/responses/utils.py | 103 +- litellm/router.py | 319 ++-- .../adaptive_router/adaptive_router.py | 10 +- litellm/router_strategy/budget_limiter.py | 12 +- .../complexity_router/complexity_router.py | 4 +- litellm/router_strategy/lowest_latency.py | 10 +- litellm/router_strategy/lowest_tpm_rpm_v2.py | 4 +- litellm/router_strategy/tag_based_routing.py | 11 +- .../add_retry_fallback_headers.py | 31 +- litellm/router_utils/batch_utils.py | 5 +- litellm/router_utils/common_utils.py | 7 +- litellm/router_utils/cooldown_cache.py | 2 +- litellm/router_utils/cooldown_callbacks.py | 7 +- litellm/router_utils/cooldown_handlers.py | 30 +- .../router_utils/fallback_event_handlers.py | 13 +- litellm/router_utils/handle_error.py | 4 +- .../router_utils/pattern_match_deployments.py | 8 +- .../encrypted_content_affinity_check.py | 4 +- .../pre_call_checks/model_rate_limit_check.py | 8 +- .../responses_api_deployment_check.py | 2 +- .../rust_bridge/callbacks_legacy_python.py | 2 +- litellm/types/agents.py | 9 +- litellm/types/completion.py | 5 +- litellm/types/integrations/prometheus.py | 14 +- litellm/types/utils.py | 11 +- litellm/utils.py | 367 +++-- .../vector_stores/vector_store_registry.py | 10 +- tests/__init__.py | 2 +- .../local_only_agent_tests/test_a2a.py | 6 +- .../test_a2a_completion_bridge.py | 10 +- tests/audio_tests/test_audio_speech.py | 6 +- tests/batches_tests/test_batch_rate_limits.py | 12 +- .../test_batches_logging_unit_tests.py | 16 +- .../test_bedrock_files_and_batches.py | 4 +- .../test_manus_files_all_methods.py | 2 +- .../test_openai_batches_and_files.py | 6 +- tests/benchmarks/test_benchmarks.py | 4 +- .../code_coverage_tests/recursive_detector.py | 8 +- .../router_code_coverage.py | 1 + .../test_javelin_guardrails.py | 2 +- tests/guardrails_tests/test_lakera_v2.py | 2 +- tests/guardrails_tests/test_presidio_pii.py | 6 +- .../test_tracing_guardrails.py | 4 +- .../base_image_generation_test.py | 2 +- tests/image_gen_tests/test_image_edits.py | 16 +- .../image_gen_tests/test_image_generation.py | 2 +- .../test_spend_log_tool_payload_content.py | 2 +- .../litellm_utils_tests/test_health_check.py | 4 +- .../test_logging_callback_manager.py | 8 +- tests/litellm_utils_tests/test_utils.py | 30 +- .../base_responses_api.py | 16 +- .../test_anthropic_responses_api.py | 4 +- .../test_azure_responses_api.py | 2 +- ...t_base_responses_api_streaming_iterator.py | 8 +- .../test_google_ai_studio_responses_api.py | 6 +- .../test_openai_responses_api.py | 38 +- .../test_responses_hooks.py | 20 +- .../base_audio_transcription_unit_tests.py | 2 +- tests/llm_translation/base_llm_unit_tests.py | 40 +- .../llm_translation/base_rerank_unit_tests.py | 2 +- .../realtime/base_realtime_tests.py | 6 +- .../test_anthropic_completion.py | 14 +- tests/llm_translation/test_azure_ai.py | 10 +- tests/llm_translation/test_azure_o_series.py | 2 +- tests/llm_translation/test_azure_openai.py | 4 +- .../llm_translation/test_bedrock_agentcore.py | 18 +- tests/llm_translation/test_bedrock_agents.py | 6 +- .../test_bedrock_completion.py | 16 +- ..._bedrock_dynamic_auth_params_unit_tests.py | 6 +- .../llm_translation/test_bedrock_embedding.py | 8 +- .../test_bedrock_invoke_tests.py | 2 +- tests/llm_translation/test_bedrock_llama.py | 2 +- .../llm_translation/test_bedrock_moonshot.py | 2 +- .../llm_translation/test_bedrock_nova_json.py | 2 +- tests/llm_translation/test_cohere.py | 2 +- .../test_deepseek_completion.py | 4 +- tests/llm_translation/test_gemini.py | 24 +- tests/llm_translation/test_gpt4o_audio.py | 2 +- .../test_litellm_proxy_provider.py | 14 +- .../test_convert_dict_to_chat_completion.py | 56 +- tests/llm_translation/test_openai.py | 10 +- tests/llm_translation/test_openrouter.py | 6 +- tests/llm_translation/test_optional_params.py | 4 +- tests/llm_translation/test_skills_api.py | 2 +- tests/local_testing/cache_unit_tests.py | 4 +- .../test_amazing_vertex_completion.py | 14 +- .../test_anthropic_prompt_caching.py | 2 +- tests/local_testing/test_blocked_user_list.py | 4 +- tests/local_testing/test_caching.py | 14 +- tests/local_testing/test_completion.py | 10 +- tests/local_testing/test_completion_cost.py | 8 +- .../test_docker_no_network_on_deploy.py | 770 ++++----- tests/local_testing/test_exceptions.py | 2 +- tests/local_testing/test_get_model_info.py | 2 +- .../test_openai_moderations_hook.py | 4 +- tests/local_testing/test_router.py | 10 +- .../test_router_batch_completion.py | 2 +- .../test_router_cooldown_handlers.py | 20 +- .../test_router_pattern_matching.py | 4 +- tests/local_testing/test_router_utils.py | 4 +- .../local_testing/test_secret_detect_hook.py | 2 +- tests/local_testing/test_streaming.py | 4 +- tests/local_testing/test_text_completion.py | 2 +- tests/logging_callback_tests/test_alerting.py | 8 +- .../test_assemble_streaming_responses.py | 16 +- .../test_bedrock_knowledgebase_hook.py | 14 +- .../test_built_in_tools_cost_tracking.py | 2 +- .../test_custom_callback_router.py | 2 +- .../test_langfuse_e2e_test.py | 8 +- .../test_langfuse_unit_tests.py | 2 +- .../test_aresponses_api_with_mcp_providers.py | 4 +- ..._anthropic_messages_prompt_caching_test.py | 12 +- ...ase_anthropic_messages_tool_search_test.py | 8 +- .../base_anthropic_unified_messages_test.py | 6 +- .../test_anthropic_messages_passthrough.py | 10 +- .../test_bedrock_anthropic_messages_test.py | 2 +- .../test_bedrock_tool_use_beta_header.py | 2 +- .../test_vertex_ai_live_passthrough.py | 6 +- .../test_websearch_interception_e2e.py | 4 +- .../proxy_admin_ui_tests/test_sso_sign_in.py | 2 +- .../test_router_batch_utils.py | 10 +- .../test_router_cooldown_per_deployment.py | 24 +- .../test_router_endpoints.py | 6 +- .../test_router_handle_error.py | 6 +- .../test_router_helper_utils.py | 22 +- tests/search_tests/base_search_unit_tests.py | 2 +- tests/search_tests/test_duckduckgo_search.py | 2 +- tests/search_tests/test_perplexity_search.py | 2 +- tests/test_litellm/__init__.py | 2 +- tests/test_litellm_rust/conftest.py | 2 +- tests/test_litellm_rust/ocr/test_lifecycle.py | 4 +- .../test_litellm_rust/ocr/test_passthrough.py | 2 +- tests/test_litellm_rust/support/isolation.py | 2 +- .../unified_google_tests/base_google_test.py | 6 +- .../base_interactions_test.py | 2 +- .../test_google_ai_studio.py | 2 +- tests/unit/batches/test_batch_utils.py | 19 +- tests/unit/caching/test_caching_handler.py | 64 +- .../unit/caching/test_redis_semantic_cache.py | 10 +- .../test_request_redis_batch_pre_call.py | 4 +- .../test_responses_stream_cache_keys.py | 8 +- tests/unit/caching/test_s3_cache.py | 2 +- tests/unit/caching/test_unit_test_caching.py | 2 +- ...responses_transformation_transformation.py | 8 +- .../test_azure_container_transformation.py | 10 +- tests/unit/containers/test_container_api.py | 10 +- .../test_container_proxy_ownership.py | 6 +- tests/unit/containers/test_container_utils.py | 2 +- .../proxy/hooks/test_managed_files.py | 164 +- .../test_afile_retrieve_returns_unified_id.py | 4 +- ...trieve_registers_missing_output_file_id.py | 4 +- ..._batch_update_db_managed_output_file_id.py | 6 +- .../test_deleted_file_returns_403_not_404.py | 14 +- .../proxy/test_file_deletion_blocking.py | 18 +- .../proxy/test_managed_files_access_check.py | 10 +- .../proxy/test_managed_files_hook.py | 32 +- .../google_genai/test_google_genai_adapter.py | 2 +- tests/unit/images/test_main.py | 2 +- .../test_slack_alerting_utils.py | 24 +- .../datadog/test_datadog_team_handler.py | 2 +- .../integrations/focus/test_export_engine.py | 70 + .../focus/test_mavvrik_destination.py | 20 +- .../gitlab/test_gitlab_prompt_manager.py | 22 +- .../test_mavvrik_focus_logger.py | 10 +- .../newrelic/test_newrelic_team_handler.py | 2 +- .../test_anthropic_cache_control_hook.py | 4 +- .../integrations/test_prometheus_labels.py | 16 +- .../llm_cost_calc/test_llm_cost_calc_utils.py | 8 +- .../test_openai_cache_write_cost.py | 2 +- .../test_responses_cache_cost_breakdown.py | 2 +- .../test_tool_call_cost_tracking.py | 22 +- .../test_convert_dict_to_response.py | 25 +- .../test_response_metadata.py | 12 +- ...llm_core_utils_prompt_templates_factory.py | 72 +- .../litellm_core_utils/test_core_helpers.py | 19 + .../litellm_core_utils/test_dd_tracing.py | 6 +- .../test_get_litellm_params.py | 19 +- .../test_health_check_helpers.py | 6 +- .../test_litellm_logging.py | 201 ++- .../test_llm_request_utils.py | 28 + .../test_model_param_helper.py | 8 +- .../test_realtime_streaming.py | 10 +- .../test_sensitive_data_masker.py | 10 +- .../test_streaming_handler.py | 10 +- .../test_streaming_overhead.py | 2 +- .../litellm_core_utils/test_token_counter.py | 8 + .../pass_through/messages/test_mcp_handler.py | 12 +- .../messages/test_streaming_iterator.py | 6 +- .../test_azure_passthrough_transformation.py | 6 +- .../llms/azure/test_azure_common_utils.py | 2 +- ...est_azure_ai_passthrough_transformation.py | 12 +- .../chat/test_streaming_choice_index.py | 226 +-- .../realtime/test_bedrock_realtime_handler.py | 2 +- ...github_copilot_responses_transformation.py | 14 +- tests/unit/llms/test_file_content_block.py | 4 +- tests/unit/llms/test_file_search_responses.py | 6 +- .../test_thought_signature_in_tool_call_id.py | 26 +- .../test_vertex_ai_gemini_transformation.py | 4 +- tests/unit/llms/vertex_ai/test_vertex.py | 32 +- .../test_xai_responses_transformation.py | 8 +- .../unit/passthrough/test_passthrough_main.py | 2 +- .../mcp_server/test_mcp_chat_completions.py | 20 +- .../proxy/batches_endpoints/test_endpoints.py | 4 +- .../test_litellm_executed_batches.py | 4 +- .../common_utils/test_check_batch_cost.py | 598 +++---- .../test_daily_spend_update_queue.py | 2 +- .../test_deferred_guardrail_logging.py | 46 +- .../guardrails/test_guardrail_coverage.py | 32 +- .../guardrails/test_guardrail_endpoints.py | 2 +- .../guardrails/test_guardrail_registry.py | 2 +- .../proxy/hooks/test_banned_keyword_list.py | 4 +- .../proxy/hooks/test_batch_file_validation.py | 46 +- .../unit/proxy/hooks/test_proxy_hooks_init.py | 4 +- .../image_endpoints/test_azure_routes.py | 2 +- .../test_model_management_endpoints.py | 2 +- .../test_files_endpoint.py | 4 +- ...test_openai_passthrough_logging_handler.py | 4 +- .../test_llm_pass_through_endpoints.py | 2 +- .../test_streaming_handler_interrupt.py | 12 +- .../proxy_server/test_background_health.py | 6 +- .../proxy/proxy_server/test_proxy_config.py | 4 +- .../proxy/proxy_server/test_routes_config.py | 4 +- .../proxy/proxy_server/test_routes_utils.py | 6 +- .../test_spend_management_endpoints.py | 2 +- .../proxy/test_common_request_processing.py | 56 +- .../unit/proxy/test_litellm_pre_call_utils.py | 4 +- tests/unit/proxy/test_prompt_test_endpoint.py | 4 +- .../unit/proxy/test_proxy_config_unit_test.py | 4 +- ...test_proxy_server_endpoints_and_startup.py | 14 +- .../proxy/test_proxy_setting_guardrails.py | 4 +- .../utils/proxy_logging/test_pre_call_hook.py | 12 +- .../test_vector_store_access_control.py | 2 +- .../test_vector_store_endpoints.py | 8 +- .../test_vector_store_rbac.py | 6 +- .../test_litellm_completion_responses.py | 32 +- .../test_streaming_iterator_transformation.py | 4 +- .../mcp/test_aresponses_api_with_mcp.py | 36 +- .../mcp/test_chat_completions_handler.py | 88 +- .../mcp/test_litellm_proxy_mcp_handler.py | 80 +- .../test_responses_prompt_management.py | 2 +- .../test_responses_router_cooldown.py | 4 +- tests/unit/responses/test_responses_utils.py | 54 +- .../test_responses_websocket_all_providers.py | 10 +- .../unit/responses/test_streaming_iterator.py | 30 +- .../test_streaming_iterator_error_events.py | 6 +- .../responses/test_text_format_conversion.py | 2 +- .../test_budget_limiter_hotpath.py | 2 +- .../test_router_routing_groups.py | 14 +- .../test_router_tag_routing.py | 44 +- .../test_encrypted_content_affinity_check.py | 90 +- .../test_responses_api_deployment_check.py | 2 +- .../unit/router_utils/test_cooldown_cache.py | 2 +- .../test_fallback_event_handlers.py | 24 +- ..._health_check_allowed_fails_integration.py | 124 +- .../test_reasoning_effort_capability.py | 30 +- .../test_router_utils_common_utils.py | 8 +- .../router_utils/test_routing_read_batch.py | 2 +- tests/unit/rust_bridge/test_logger.py | 8 +- tests/unit/test_cost_calculator.py | 46 +- tests/unit/test_deepseek_model_metadata.py | 6 +- tests/unit/test_logging.py | 6 +- tests/unit/test_main.py | 60 +- tests/unit/test_model_param_helper.py | 6 +- tests/unit/test_private_usage_aliases.py | 1052 ++++++++++++ .../unit/test_redact_string_in_error_paths.py | 26 +- tests/unit/test_redis.py | 8 +- .../test_register_model_custom_pricing.py | 19 +- .../test_responses_api_bridge_non_stream.py | 16 +- tests/unit/test_router/test_router.py | 1421 +++++++---------- tests/unit/test_router_block_helpers.py | 6 +- .../unit/test_router_model_cost_isolation.py | 2 +- tests/unit/test_router_order_fallback.py | 22 +- tests/unit/test_router_weighted_failover.py | 6 +- tests/unit/test_secret_redaction.py | 5 +- tests/unit/test_utils.py | 29 +- tests/unit/types/test_completion.py | 6 +- .../test_prometheus_label_value_sanitize.py | 4 +- .../test_vector_store_registry.py | 8 +- .../base_vector_store_test.py | 4 +- .../vector_store_tests/rag/base_rag_tests.py | 4 +- .../vector_store_tests/rag/test_rag_openai.py | 4 +- .../rag/test_rag_vertex_ai.py | 4 +- .../test_azure_ai_vector_store.py | 2 +- .../test_milvus_vector_store.py | 2 +- .../test_ragflow_vector_store.py | 2 +- 461 files changed, 6719 insertions(+), 4996 deletions(-) create mode 100644 tests/unit/integrations/focus/test_export_engine.py create mode 100644 tests/unit/test_private_usage_aliases.py diff --git a/enterprise/enterprise_hooks/__init__.py b/enterprise/enterprise_hooks/__init__.py index e93c8c9150a..6de3325f113 100644 --- a/enterprise/enterprise_hooks/__init__.py +++ b/enterprise/enterprise_hooks/__init__.py @@ -1,15 +1,15 @@ from typing import Dict, Literal, Type, Union -from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles +from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles from litellm_enterprise.proxy.hooks.managed_vector_stores import ( - _PROXY_LiteLLMManagedVectorStores, + PROXY_LiteLLMManagedVectorStores, ) from litellm.integrations.custom_logger import CustomLogger ENTERPRISE_PROXY_HOOKS: Dict[str, Type[CustomLogger]] = { - "managed_files": _PROXY_LiteLLMManagedFiles, - "managed_vector_stores": _PROXY_LiteLLMManagedVectorStores, + "managed_files": PROXY_LiteLLMManagedFiles, + "managed_vector_stores": PROXY_LiteLLMManagedVectorStores, } diff --git a/enterprise/enterprise_hooks/banned_keywords.py b/enterprise/enterprise_hooks/banned_keywords.py index 6f6a37b6c55..97315e88732 100644 --- a/enterprise/enterprise_hooks/banned_keywords.py +++ b/enterprise/enterprise_hooks/banned_keywords.py @@ -20,7 +20,7 @@ from litellm._logging import verbose_proxy_logger from fastapi import HTTPException -class _ENTERPRISE_BannedKeywords(CustomLogger): +class ENTERPRISE_BannedKeywords(CustomLogger): enforces_request_content: bool = True # Class variables or attributes def __init__(self): @@ -114,3 +114,4 @@ class _ENTERPRISE_BannedKeywords(CustomLogger): response: str, ): self.test_violation(test_str=response) +_ENTERPRISE_BannedKeywords = ENTERPRISE_BannedKeywords diff --git a/enterprise/enterprise_hooks/blocked_user_list.py b/enterprise/enterprise_hooks/blocked_user_list.py index dfaf91ea081..f6998d58cd5 100644 --- a/enterprise/enterprise_hooks/blocked_user_list.py +++ b/enterprise/enterprise_hooks/blocked_user_list.py @@ -21,7 +21,7 @@ from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET from litellm.proxy.utils import PrismaClient -class _ENTERPRISE_BlockedUserList(CustomLogger): +class ENTERPRISE_BlockedUserList(CustomLogger): enforces_request_content: bool = True # Class variables or attributes def __init__(self, prisma_client: Optional[PrismaClient]): @@ -128,3 +128,4 @@ class _ENTERPRISE_BlockedUserList(CustomLogger): str(e) ) ) +_ENTERPRISE_BlockedUserList = ENTERPRISE_BlockedUserList diff --git a/enterprise/enterprise_hooks/google_text_moderation.py b/enterprise/enterprise_hooks/google_text_moderation.py index 5b2d71c5cca..b3ef5a3d496 100644 --- a/enterprise/enterprise_hooks/google_text_moderation.py +++ b/enterprise/enterprise_hooks/google_text_moderation.py @@ -16,7 +16,7 @@ from litellm.proxy.guardrails._content_utils import iter_message_text from litellm.types.utils import CallTypesLiteral -class _ENTERPRISE_GoogleTextModeration(CustomLogger): +class ENTERPRISE_GoogleTextModeration(CustomLogger): user_api_key_cache = None confidence_categories = [ "toxic", @@ -125,9 +125,10 @@ class _ENTERPRISE_GoogleTextModeration(CustomLogger): ) # Handle the response return data +_ENTERPRISE_GoogleTextModeration = ENTERPRISE_GoogleTextModeration -# google_text_moderation_obj = _ENTERPRISE_GoogleTextModeration() +# google_text_moderation_obj = ENTERPRISE_GoogleTextModeration() # asyncio.run( # google_text_moderation_obj.async_moderation_hook( # data={"messages": [{"role": "user", "content": "Hey, how's it going?"}]} diff --git a/enterprise/enterprise_hooks/openai_moderation.py b/enterprise/enterprise_hooks/openai_moderation.py index 017f51bfabd..d8536ad1ac3 100644 --- a/enterprise/enterprise_hooks/openai_moderation.py +++ b/enterprise/enterprise_hooks/openai_moderation.py @@ -24,7 +24,7 @@ from litellm.proxy.guardrails._content_utils import iter_message_text from litellm.types.utils import CallTypesLiteral -class _ENTERPRISE_OpenAI_Moderation(CustomLogger): +class ENTERPRISE_OpenAI_Moderation(CustomLogger): @property def model_name(self) -> str: return litellm.openai_moderations_model_name or DEFAULT_OPENAI_MODERATIONS_MODEL @@ -55,3 +55,4 @@ class _ENTERPRISE_OpenAI_Moderation(CustomLogger): status_code=403, detail={"error": "Violated content safety policy"} ) pass +_ENTERPRISE_OpenAI_Moderation = ENTERPRISE_OpenAI_Moderation diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py index cf22488edcb..36e2bf865b3 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py @@ -3,9 +3,10 @@ Endpoints for managing email alerts on litellm """ import json -from typing import Dict +from typing import Dict, Final, cast from fastapi import APIRouter, Depends, HTTPException +from pydantic import JsonValue from litellm_enterprise.types.enterprise_callbacks.send_emails import ( DefaultEmailSettings, EmailEvent, @@ -38,14 +39,15 @@ async def _get_email_settings(prisma_client) -> Dict[str, bool]: and general_settings_entry.param_value is not None ): # Get general settings value - if isinstance(general_settings_entry.param_value, str): - general_settings = json.loads(general_settings_entry.param_value) - else: - general_settings = general_settings_entry.param_value + general_settings: Final = ( + cast(Dict[str, object], json.loads(general_settings_entry.param_value)) + if isinstance(general_settings_entry.param_value, str) + else cast(Dict[str, object], general_settings_entry.param_value) + ) # Extract email_settings from general settings if it exists if general_settings and "email_settings" in general_settings: - email_settings = general_settings["email_settings"] + email_settings: Final = cast(Dict[str, bool], general_settings["email_settings"]) # Update settings_dict with values from general_settings for event_name, enabled in email_settings.items(): settings_dict[event_name] = enabled @@ -64,7 +66,7 @@ async def _save_email_settings(prisma_client, settings: Dict[str, bool]): from litellm.proxy.proxy_server import proxy_config proxy_config.reject_config_owned_writes( - section_name="general_settings", changed_keys={"email_settings": settings} + section_name="general_settings", changed_keys={"email_settings": cast(JsonValue, settings)} ) try: verbose_proxy_logger.debug( @@ -77,16 +79,15 @@ async def _save_email_settings(prisma_client, settings: Dict[str, bool]): ) # Initialize general settings dict - if ( - general_settings_entry is not None - and general_settings_entry.param_value is not None - ): - if isinstance(general_settings_entry.param_value, str): - general_settings = json.loads(general_settings_entry.param_value) - else: - general_settings = dict(general_settings_entry.param_value) - else: - general_settings = {} + general_settings: Final = ( + ( + cast(Dict[str, object], json.loads(general_settings_entry.param_value)) + if isinstance(general_settings_entry.param_value, str) + else cast(Dict[str, object], dict(general_settings_entry.param_value)) + ) + if general_settings_entry is not None and general_settings_entry.param_value is not None + else {} + ) # Update email_settings in general_settings general_settings["email_settings"] = settings diff --git a/enterprise/litellm_enterprise/integrations/custom_guardrail.py b/enterprise/litellm_enterprise/integrations/custom_guardrail.py index f07752d5c18..cbac4cb5b15 100644 --- a/enterprise/litellm_enterprise/integrations/custom_guardrail.py +++ b/enterprise/litellm_enterprise/integrations/custom_guardrail.py @@ -36,7 +36,7 @@ class EnterpriseCustomGuardrailHelper: proxy_server_request = data.get("proxy_server_request", {}) - request_tags = StandardLoggingPayloadSetup._get_request_tags( + request_tags = StandardLoggingPayloadSetup.get_request_tags( litellm_params=data, proxy_server_request=proxy_server_request, ) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 985a9b6de38..84e319c8c37 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -60,6 +60,32 @@ class _ManagedObjectRow(Protocol): @property def file_object(self) -> object: ... + @property + def org_id(self) -> str | None: ... + + @property + def api_key(self) -> str | None: ... + + @property + def team_id(self) -> str | None: ... + + +class _ReadableFileContent(Protocol): + async def read(self) -> bytes: ... + + +class _HasFileContent(Protocol): + @property + def content(self) -> bytes: ... + + +async def _file_content_bytes(file_content: object) -> bytes: + if hasattr(file_content, "content"): + return cast(_HasFileContent, file_content).content + if hasattr(file_content, "read"): + return await cast(_ReadableFileContent, file_content).read() + return cast(bytes, file_content) + def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": return ManagedObjectRepository(prisma_client).table @@ -175,17 +201,21 @@ class CheckBatchCost: return None async def _get_org_id(self, job: "_ManagedObjectRow", batch_id: str) -> str | None: - org_id = getattr(job, "org_id", None) + org_id: Final = getattr(job, "org_id", None) if org_id: return org_id - api_key = getattr(job, "api_key", None) - team_id = getattr(job, "team_id", None) + api_key: Final = getattr(job, "api_key", None) + team_id: Final = getattr(job, "team_id", None) if api_key: try: key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table( self.prisma_client ).find_unique(where={"token": api_key}) - key_org_id = getattr(key_row, "organization_id", None) if key_row is not None else None + key_org_id: Final = ( + cast(str | None, getattr(key_row, "organization_id", None)) + if key_row is not None + else None + ) if key_org_id: return key_org_id except Exception as e: @@ -199,7 +229,7 @@ class CheckBatchCost: team_row: prisma_models.LiteLLM_TeamTable | None = await _team_table(self.prisma_client).find_unique( where={"team_id": team_id} ) - return getattr(team_row, "organization_id", None) if team_row is not None else None + return cast(str | None, getattr(team_row, "organization_id", None)) if team_row is not None else None except Exception as e: verbose_proxy_logger.error(f"CheckBatchCost: could not resolve the team's org for batch {batch_id}: {e}") return None @@ -673,16 +703,17 @@ class CheckBatchCost: from litellm.types.utils import LiteLLMBatch - file_object = job.file_object - if isinstance(file_object, str): - try: - file_object = json.loads(file_object) - except (json.JSONDecodeError, ValueError): - return None - if not isinstance(file_object, dict): + file_object: Final = job.file_object + try: + parsed_file_object: Final[object] = ( + json.loads(file_object) if isinstance(file_object, str) else file_object + ) + except (json.JSONDecodeError, ValueError): + return None + if not isinstance(parsed_file_object, dict): return None try: - return LiteLLMBatch.model_validate(file_object).input_file_id + return LiteLLMBatch.model_validate(parsed_file_object).input_file_id except Exception: return None @@ -704,10 +735,11 @@ class CheckBatchCost: """ from litellm.batches.batch_utils import ( count_error_file_failed_requests, - _get_file_content_as_dictionary, + get_file_content_as_dictionary, calculate_batch_cost_and_usage, ) from litellm.files.main import afile_content + from litellm.files.types import FileContentCallOptions from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.litellm_logging import deployment_pricing_model_info, mask_api_base_credentials @@ -745,30 +777,22 @@ class CheckBatchCost: except (IndexError, AttributeError): pass - credentials = self.llm_router.get_deployment_credentials_with_provider(model_id) or {} + credentials: Final = self.llm_router.get_deployment_credentials_with_provider(model_id) or {} _file_content = await afile_content( - file_id=raw_output_file_id, + file_id=cast(str, raw_output_file_id), _litellm_internal_model_credentials=MappingProxyType(dict(credentials)), - **credentials, + **cast(FileContentCallOptions, credentials), ) - # Access content - handle both direct attribute and method call - if hasattr(_file_content, 'content'): - content_bytes = _file_content.content # type: ignore[union-attr] - elif hasattr(_file_content, 'read'): - content_bytes = await _file_content.read() # type: ignore[misc] - else: - content_bytes = _file_content # type: ignore[assignment] + content_bytes: Final = await _file_content_bytes(_file_content) - file_content_as_dict = _get_file_content_as_dictionary( - content_bytes # type: ignore[arg-type] - ) + file_content_as_dict = get_file_content_as_dictionary(content_bytes) # Record output file size if prom_logger and content_bytes: try: prom_logger.record_managed_file_size( - size_bytes=len(content_bytes), # type: ignore + size_bytes=len(content_bytes), purpose="batch", file_type="output", model=model_id, @@ -805,7 +829,7 @@ class CheckBatchCost: team_id=getattr(job, "team_id", None), ) for _file_attr in ["output_file_id", "error_file_id"]: - _raw_file_id = getattr(response, _file_attr, None) + _raw_file_id = cast(str | None, getattr(response, _file_attr, None)) if _raw_file_id and not _is_base64_encoded_unified_file_id(_raw_file_id): try: _unified_file_id = managed_files_hook.get_unified_output_file_id( diff --git a/enterprise/litellm_enterprise/proxy/enterprise_routes.py b/enterprise/litellm_enterprise/proxy/enterprise_routes.py index a76b8d01f0e..df3e5470d46 100644 --- a/enterprise/litellm_enterprise/proxy/enterprise_routes.py +++ b/enterprise/litellm_enterprise/proxy/enterprise_routes.py @@ -8,7 +8,7 @@ from . import ui_crud_endpoints # side-effect: registers extra UI settings from .audit_logging_endpoints import router as audit_logging_router from .liteadmin import router as liteadmin_router from .management_endpoints import management_endpoints_router -from .utils import _should_block_robots +from .utils import should_block_robots __all__ = ["router", "ui_crud_endpoints"] @@ -25,7 +25,7 @@ async def get_robots(): Block all web crawlers from indexing the proxy server endpoints This is useful for ensuring that the API endpoints aren't indexed by search engines """ - if _should_block_robots(): + if should_block_robots(): return Response(content="User-agent: *\nDisallow: /", media_type="text/plain") else: return Response(status_code=404) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 011bf95defd..aae49402146 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -23,6 +23,7 @@ from uuid import NAMESPACE_URL, uuid5 import httpx from fastapi import HTTPException from pydantic import ValidationError +from typing_extensions import ReadOnly import litellm from litellm import Router, verbose_logger @@ -30,10 +31,12 @@ from litellm._internal_context import with_service_target from litellm._uuid import uuid from litellm.caching.caching import DualCache from litellm.constants import MAX_FILE_LIST_LIMIT +from litellm.files.types import FileRetrieveCallOptions, FileRetrieveProvider from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.prompt_templates.common_utils import ( extract_file_metadata, ) +from openai import AsyncOpenAI from openai.types.file_deleted import FileDeleted from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend @@ -190,6 +193,19 @@ class _ManagedObjectTableActions(Protocol): async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... +class _ManagedResourceDatabase(Protocol): + @property + def litellm_managedfiletable(self) -> _ManagedFileTableActions: ... + + @property + def litellm_managedobjecttable(self) -> _ManagedObjectTableActions: ... + + +class _ManagedResourcePrismaClient(Protocol): + @property + def db(self) -> _ManagedResourceDatabase: ... + + class _SchedulerWithJobLookup(Protocol): def get_job(self, job_id: str) -> object: ... @@ -199,7 +215,12 @@ class _CursorPageArgs(TypedDict, total=False): skip: int -def _managed_file_table(prisma_client: PrismaClient) -> _ManagedFileTableActions: +class _RouterFileCallKwargs(TypedDict, total=False): + client: ReadOnly[AsyncOpenAI | None] + custom_llm_provider: ReadOnly[str | None] + + +def _managed_file_table(prisma_client: _ManagedResourcePrismaClient) -> _ManagedFileTableActions: return prisma_client.db.litellm_managedfiletable @@ -213,7 +234,7 @@ def _iter_provider_file_id_pairs( yield provider_file_id, row.unified_file_id -def _managed_object_table(prisma_client: PrismaClient) -> _ManagedObjectTableActions: +def _managed_object_table(prisma_client: _ManagedResourcePrismaClient) -> _ManagedObjectTableActions: return prisma_client.db.litellm_managedobjecttable @@ -233,7 +254,7 @@ def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, s _MANAGED_FILES_TARGET: Final = "managed_files" -class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): +class PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # Class variables or attributes def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient): self.internal_usage_cache = internal_usage_cache @@ -1282,7 +1303,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): target_model_names_list=target_model_names_list, litellm_parent_otel_span=litellm_parent_otel_span, ) - response = await _PROXY_LiteLLMManagedFiles.return_unified_file_id( + response = await PROXY_LiteLLMManagedFiles.return_unified_file_id( file_objects=responses, create_file_request=create_file_request, internal_usage_cache=self.internal_usage_cache, @@ -1459,13 +1480,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): _creds = _llm_router.get_deployment_credentials_with_provider(model_id) or {} file_object = await litellm.afile_retrieve( file_id=provider_file_id, - **_creds, + **cast(FileRetrieveCallOptions, _creds), ) else: file_object = await litellm.afile_retrieve( - custom_llm_provider=model_name.split("/")[0] - if model_name and "/" in model_name - else "openai", # type: ignore[arg-type] + custom_llm_provider=cast( + FileRetrieveProvider, + model_name.split("/")[0] if model_name and "/" in model_name else "openai", + ), file_id=provider_file_id, ) verbose_logger.debug( @@ -1608,8 +1630,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): try: model_id, model_file_id = next(iter(stored_file_object.model_mappings.items())) - credentials = llm_router.get_deployment_credentials_with_provider(model_id) or {} - response = await litellm.afile_retrieve(file_id=model_file_id, **credentials) + credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id) or {} + response = await litellm.afile_retrieve( + file_id=model_file_id, + **cast(FileRetrieveCallOptions, credentials), + ) response.id = file_id # Replace with unified ID return response except Exception as e: @@ -1904,7 +1929,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): else {} ), } - await llm_router.afile_delete(model=model_id, file_id=model_file_id, **delete_data) + await llm_router.afile_delete( + model=model_id, + file_id=model_file_id, + **cast(_RouterFileCallKwargs, delete_data), + ) async def afile_content( self, @@ -1940,7 +1969,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): data["_litellm_internal_model_credentials"] = cast(Dict, MappingProxyType(dict(credentials))) else: data.pop("_litellm_internal_model_credentials", None) - return await llm_router.afile_content(model=model_id, file_id=provider_file_id, **data) # type: ignore + return cast( + HttpxBinaryResponseContent, + await llm_router.afile_content( + model=model_id, + file_id=provider_file_id, + **cast(_RouterFileCallKwargs, data), + ), + ) except Exception as e: exception_dict[model_id] = str(e) raise Exception( @@ -2068,3 +2104,4 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): verbose_logger.debug( f"Converted file {file_id} from storage backend to base64 with format {content_type}" ) +_PROXY_LiteLLMManagedFiles = PROXY_LiteLLMManagedFiles diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py index 3b8c19f0097..c53f359a468 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py @@ -2,11 +2,11 @@ ## This hook is used to manage vector stores with target_model_names support ## It allows creating vector stores across multiple models and managing them with unified IDs -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast +from collections.abc import Mapping +from typing import TYPE_CHECKING, Any, Dict, Final, List, Optional, Union, cast from fastapi import HTTPException -import litellm from litellm import Router, verbose_logger from litellm._uuid import uuid from litellm.integrations.custom_logger import CustomLogger @@ -38,7 +38,7 @@ else: PrismaClient = Any -class _PROXY_LiteLLMManagedVectorStores( +class PROXY_LiteLLMManagedVectorStores( CustomLogger, BaseManagedResource[VectorStoreCreateResponse] ): """ @@ -89,10 +89,10 @@ class _PROXY_LiteLLMManagedVectorStores( # Model ID is stored in hidden params if the response object supports it # For TypedDict responses, we need to check if _hidden_params was added - hidden_params: Dict[str, Any] = {} + hidden_params: Mapping[str, object] = {} if hasattr(resource_object, "_hidden_params"): - hidden_params = getattr(resource_object, "_hidden_params", {}) or {} - model_id = hidden_params.get("model_id", "") + hidden_params = cast(Mapping[str, object], getattr(resource_object, "_hidden_params", {}) or {}) + model_id: Final = cast(str, hidden_params.get("model_id", "")) return generate_unified_id_string( resource_type=self.resource_type, @@ -106,7 +106,7 @@ class _PROXY_LiteLLMManagedVectorStores( self, llm_router: Router, model: str, - request_data: Dict[str, Any], + request_data: dict[str, object] | VectorStoreCreateOptionalRequestParams, litellm_parent_otel_span: Span, ) -> VectorStoreCreateResponse: """ @@ -122,10 +122,8 @@ class _PROXY_LiteLLMManagedVectorStores( VectorStoreCreateResponse from the provider """ # Use the router to create the vector store - response = await llm_router.avector_store_create( - model=model, **request_data - ) - return response + response: Final = await llm_router.avector_store_create(model=model, **request_data) + return cast(VectorStoreCreateResponse, response) # ============================================================================ # VECTOR STORE CRUD OPERATIONS @@ -464,3 +462,4 @@ class _PROXY_LiteLLMManagedVectorStores( parent_otel_span=parent_otel_span, resource_id_key="vector_store_id", ) +_PROXY_LiteLLMManagedVectorStores = PROXY_LiteLLMManagedVectorStores diff --git a/enterprise/litellm_enterprise/proxy/proxy_server.py b/enterprise/litellm_enterprise/proxy/proxy_server.py index 79d3ebdf9ee..8e90ef62054 100644 --- a/enterprise/litellm_enterprise/proxy/proxy_server.py +++ b/enterprise/litellm_enterprise/proxy/proxy_server.py @@ -9,7 +9,7 @@ custom_auth_settings: Optional[CustomAuthSettings] = None class EnterpriseProxyConfig: async def load_custom_auth_settings( self, general_settings: dict - ) -> CustomAuthSettings: + ) -> Optional[CustomAuthSettings]: custom_auth_settings = general_settings.get("custom_auth_settings", None) if custom_auth_settings is not None: custom_auth_settings = CustomAuthSettings( diff --git a/enterprise/litellm_enterprise/proxy/utils.py b/enterprise/litellm_enterprise/proxy/utils.py index 227ea0a9ff0..50d9b658afa 100644 --- a/enterprise/litellm_enterprise/proxy/utils.py +++ b/enterprise/litellm_enterprise/proxy/utils.py @@ -3,7 +3,7 @@ from typing import Optional, Union from litellm.secret_managers.main import str_to_bool -def _should_block_robots(): +def should_block_robots(): """ Returns True if the robots.txt file should block web crawlers @@ -33,3 +33,4 @@ def _should_block_robots(): ) return True return False +_should_block_robots = should_block_robots diff --git a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py index e95a7c99971..9b87f4f1a7f 100644 --- a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py +++ b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py @@ -171,7 +171,7 @@ async def list_vector_stores( try: # Get vector stores from database (source of truth) # Only return what's in the database to ensure consistency across instances - vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db( + vector_stores_from_db = await VectorStoreRegistry.get_vector_stores_from_db( prisma_client=prisma_client ) diff --git a/litellm/__init__.py b/litellm/__init__.py index 487d5da1663..9cfaccf00e9 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -53,10 +53,10 @@ from litellm.types.integrations.pointfive import PointFiveInitParams from litellm.types.integrations.zerobus import ZerobusInitParams from litellm._logging import ( set_verbose, - _turn_on_debug, + turn_on_debug, verbose_logger, json_logs, - _turn_on_json, + turn_on_json, log_level, ) import re @@ -108,10 +108,10 @@ litellm_mode = os.getenv("LITELLM_MODE", "DEV") # "PRODUCTION", "DEV" #################################################### if set_verbose: - _turn_on_debug() + turn_on_debug() #################################################### ### Callbacks /Logging / Success / Failure Handlers ##### -CALLBACK_TYPES = Union[str, Callable, "CustomLogger"] # CustomLogger is lazy-loaded +CALLBACK_TYPES = Union[str, Callable[..., object], "CustomLogger"] # CustomLogger is lazy-loaded input_callback: List[CALLBACK_TYPES] = [] success_callback: List[CALLBACK_TYPES] = [] failure_callback: List[CALLBACK_TYPES] = [] @@ -177,9 +177,9 @@ _custom_logger_compatible_callbacks_literal = Literal[ ] 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)) +_known_custom_logger_compatible_callbacks: List[str] = list(get_args(_custom_logger_compatible_callbacks_literal)) callbacks: List[ - Union[Callable, _custom_logger_compatible_callbacks_literal, "CustomLogger"] # CustomLogger is lazy-loaded + Union[Callable[..., object], str, "CustomLogger"] # CustomLogger is lazy-loaded ] = [] callback_settings: Dict[str, Dict[str, Any]] = {} initialized_langfuse_clients: int = 0 @@ -194,13 +194,13 @@ datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged p 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 +_async_input_callback: List[Union[str, Callable[..., object], "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 +_async_success_callback: List[Union[str, Callable[..., object], "CustomLogger"]] = ( # CustomLogger is lazy-loaded [] ) # internal variable - async custom callbacks are routed here. -_async_failure_callback: List[Union[str, Callable, "CustomLogger"]] = ( # CustomLogger is lazy-loaded +_async_failure_callback: List[Union[str, Callable[..., object], "CustomLogger"]] = ( # CustomLogger is lazy-loaded [] ) # internal variable - async custom callbacks are routed here. pre_call_rules: List[Callable] = [] @@ -1486,6 +1486,7 @@ from .realtime_api.main import ( arealtime_calls, ) from .responses.main import _aresponses_websocket + from .fine_tuning.main import * from .files.main import * from .vector_store_files.main import ( @@ -1527,6 +1528,9 @@ from . import rag ### CUSTOM LLMs ### from .types.llms.custom_llm import CustomLLMItem +_turn_on_debug = turn_on_debug +_turn_on_json = turn_on_json + 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] = ( @@ -2240,7 +2244,9 @@ if TYPE_CHECKING: register_model: Callable[..., None] encode: Callable[..., list] decode: Callable[..., str] + calculate_retry_after: Callable[..., float] _calculate_retry_after: Callable[..., float] + should_retry: Callable[..., bool] _should_retry: Callable[..., bool] get_supported_openai_params: Callable[..., Optional[list]] get_api_base: Callable[..., Optional[str]] @@ -2336,9 +2342,9 @@ def __getattr__(name: str) -> Any: _async_client_cleanup_registered = True # Use cached registry from _lazy_imports instead of importing tuples every time - from ._lazy_imports import _get_lazy_import_registry + from ._lazy_imports import get_lazy_import_registry - registry: Final = _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: diff --git a/litellm/_lazy_imports.py b/litellm/_lazy_imports.py index d4a12e7c1e5..46126c444a6 100644 --- a/litellm/_lazy_imports.py +++ b/litellm/_lazy_imports.py @@ -91,12 +91,15 @@ def _get_module_level_client_timeout(litellm_globals: Mapping[str, Any]) -> "flo # They're separate from the main lazy import system because they have specific use cases -def _get_default_encoding() -> "Tokenizer": +def get_default_encoding() -> "Tokenizer": from litellm.rust_bridge.tokenizer import get_encoding return get_encoding("cl100k_base") +_get_default_encoding = get_default_encoding + + # Lazy loader for get_modified_max_tokens to avoid importing token_counter at module import time _get_modified_max_tokens_func: "Callable[..., int | None] | None" = None @@ -126,7 +129,7 @@ _token_counter_new_func: "Callable[..., int] | None" = None _messages_reach_token_count_func: "Callable[..., bool] | None" = None -def _get_token_counter_new() -> "Callable[..., int]": +def get_token_counter_new() -> "Callable[..., int]": """ Lazily load and cache the token_counter function (aliased as token_counter_new). @@ -146,8 +149,11 @@ def _get_token_counter_new() -> "Callable[..., int]": return _token_counter_new_func -def _get_messages_reach_token_count() -> "Callable[..., bool]": - """Lazily load ``messages_reach_token_count`` for the same reason as ``_get_token_counter_new``.""" +_get_token_counter_new = get_token_counter_new + + +def get_messages_reach_token_count() -> "Callable[..., bool]": + """Lazily load ``messages_reach_token_count`` for the same reason as ``get_token_counter_new``.""" global _messages_reach_token_count_func if _messages_reach_token_count_func is None: from litellm.litellm_core_utils.token_counter import ( @@ -158,6 +164,9 @@ def _get_messages_reach_token_count() -> "Callable[..., bool]": return _messages_reach_token_count_func +_get_messages_reach_token_count = get_messages_reach_token_count + + # ============================================================================ # MAIN LAZY IMPORT SYSTEM # ============================================================================ @@ -168,7 +177,7 @@ def _get_messages_reach_token_count() -> "Callable[..., bool]": _LAZY_IMPORT_REGISTRY: dict[str, Callable[[str], object]] | None = None -def _get_lazy_import_registry() -> dict[str, Callable[[str], object]]: +def get_lazy_import_registry() -> dict[str, Callable[[str], object]]: """ Build the registry that maps attribute names to their handler functions. @@ -217,6 +226,9 @@ def _get_lazy_import_registry() -> dict[str, Callable[[str], object]]: return _LAZY_IMPORT_REGISTRY +_get_lazy_import_registry = get_lazy_import_registry + + class _AttributeView(TypedDict): """Holds one module attribute so the lazily fetched value is read back as ``object``.""" diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index fcd2eed5387..78f5e4ef04b 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -48,7 +48,9 @@ UTILS_NAMES: Final = ( "register_model", "encode", "decode", + "calculate_retry_after", "_calculate_retry_after", + "should_retry", "_should_retry", "get_supported_openai_params", "get_api_base", @@ -384,15 +386,18 @@ UTILS_MODULE_NAMES: Final = ( "_get_response_headers", "get_llm_provider", "_is_non_openai_azure_model", + "is_non_openai_azure_model", "get_supported_openai_params", "LiteLLMResponseObjectHandler", "_handle_invalid_parallel_tool_calls", + "handle_invalid_parallel_tool_calls", "convert_to_model_response_object", "convert_to_streaming_response", "convert_to_streaming_response_async", "get_api_base", "ResponseMetadata", "_parse_content_for_reasoning", + "parse_content_for_reasoning", "LiteLLMLoggingObject", "redact_message_input_output_from_logging", "CustomStreamWrapper", @@ -415,8 +420,10 @@ UTILS_MODULE_NAMES: Final = ( "delete_nested_value", "is_nested_path", "_get_base_model_from_litellm_call_metadata", + "get_base_model_from_litellm_call_metadata", "get_litellm_params", "_ensure_extra_body_is_safe", + "ensure_extra_body_is_safe", "get_formatted_prompt", "get_response_headers", "update_response_metadata", @@ -475,8 +482,10 @@ _UTILS_IMPORT_MAP: Final = { "register_model": (".utils", "register_model"), "encode": (".utils", "encode"), "decode": (".utils", "decode"), - "_calculate_retry_after": (".utils", "_calculate_retry_after"), - "_should_retry": (".utils", "_should_retry"), + "calculate_retry_after": (".utils", "calculate_retry_after"), + "_calculate_retry_after": (".utils", "calculate_retry_after"), + "should_retry": (".utils", "should_retry"), + "_should_retry": (".utils", "should_retry"), "get_supported_openai_params": (".utils", "get_supported_openai_params"), "get_api_base": (".utils", "get_api_base"), "get_first_chars_messages": (".utils", "get_first_chars_messages"), @@ -1320,7 +1329,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = { ), "_is_non_openai_azure_model": ( "litellm.litellm_core_utils.get_llm_provider_logic", - "_is_non_openai_azure_model", + "is_non_openai_azure_model", + ), + "is_non_openai_azure_model": ( + "litellm.litellm_core_utils.get_llm_provider_logic", + "is_non_openai_azure_model", ), "get_supported_openai_params": ( "litellm.litellm_core_utils.get_supported_openai_params", @@ -1332,7 +1345,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = { ), "_handle_invalid_parallel_tool_calls": ( "litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response", - "_handle_invalid_parallel_tool_calls", + "handle_invalid_parallel_tool_calls", + ), + "handle_invalid_parallel_tool_calls": ( + "litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response", + "handle_invalid_parallel_tool_calls", ), "convert_to_model_response_object": ( "litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response", @@ -1356,7 +1373,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = { ), "_parse_content_for_reasoning": ( "litellm.litellm_core_utils.prompt_templates.common_utils", - "_parse_content_for_reasoning", + "parse_content_for_reasoning", + ), + "parse_content_for_reasoning": ( + "litellm.litellm_core_utils.prompt_templates.common_utils", + "parse_content_for_reasoning", ), "LiteLLMLoggingObject": ( "litellm.litellm_core_utils.redact_messages", @@ -1429,7 +1450,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = { ), "_get_base_model_from_litellm_call_metadata": ( "litellm.litellm_core_utils.get_litellm_params", - "_get_base_model_from_litellm_call_metadata", + "get_base_model_from_litellm_call_metadata", + ), + "get_base_model_from_litellm_call_metadata": ( + "litellm.litellm_core_utils.get_litellm_params", + "get_base_model_from_litellm_call_metadata", ), "get_litellm_params": ( "litellm.litellm_core_utils.get_litellm_params", @@ -1437,7 +1462,11 @@ _UTILS_MODULE_IMPORT_MAP: Final = { ), "_ensure_extra_body_is_safe": ( "litellm.litellm_core_utils.llm_request_utils", - "_ensure_extra_body_is_safe", + "ensure_extra_body_is_safe", + ), + "ensure_extra_body_is_safe": ( + "litellm.litellm_core_utils.llm_request_utils", + "ensure_extra_body_is_safe", ), "get_formatted_prompt": ( "litellm.litellm_core_utils.llm_response_utils.get_formatted_prompt", diff --git a/litellm/_logging.py b/litellm/_logging.py index 5e02ff8de35..2b5c5dd548f 100644 --- a/litellm/_logging.py +++ b/litellm/_logging.py @@ -23,10 +23,12 @@ from litellm.litellm_core_utils.env_utils import get_env_int from litellm.litellm_core_utils.safe_json_dumps import UNSERIALIZABLE_OBJECT, safe_dumps, safe_json_structure from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.litellm_core_utils.secret_redaction import ( - _python_redact_string, - _python_redact_structured_value, + python_redact_string, + python_redact_structured_value, redact_internal_details, - redact_string, +) +from litellm.litellm_core_utils.secret_redaction import ( + redact_string as redact_secret_string, ) from litellm.rust_bridge import diagnostics @@ -52,7 +54,7 @@ def _sanitize_correlation_id(value: str) -> str: pass through credential redaction. """ stripped: Final = "".join(ch for ch in value if ch.isprintable()) - return _redact_string(stripped)[:_MAX_CORRELATION_ID_LENGTH] + return redact_string(stripped)[:_MAX_CORRELATION_ID_LENGTH] def set_session_id(session_id: str) -> "contextvars.Token[str]": @@ -71,10 +73,13 @@ if set_verbose is True: _ENABLE_SECRET_REDACTION: Final = os.getenv("LITELLM_DISABLE_REDACT_SECRETS", "").lower() != "true" -def _redact_string(value: str) -> str: +def redact_string(value: str) -> str: if not _ENABLE_SECRET_REDACTION: return value - return redact_string(value) + return redact_secret_string(value) + + +_redact_string = redact_string _REDACTED_RECORD_ATTR: Final = "litellm_redacted" @@ -114,7 +119,7 @@ def redact_secrets(value: str) -> str: """ if not _ENABLE_SECRET_REDACTION: return value - return _redact_string(value) + return redact_string(value) def redact_internal_details_from_client_message(value: str) -> str: @@ -165,7 +170,7 @@ _REDACTION_PLACEHOLDER: Final = "REDACTED" def _hides_a_credential(value: str) -> bool: """Whether *value* only looks clean until it is percent-decoded.""" decoded: Final = unquote(value) - return _python_redact_string(decoded) != decoded + return python_redact_string(decoded) != decoded def _drop_encoded_credential(scrubbed: str) -> str: @@ -193,10 +198,10 @@ def _scrub_access_arg(value: str) -> str: pattern and would then be logged raw. """ if len(value) <= _MAX_SCRUBBED_ACCESS_ARG: - return _drop_encoded_credential(_python_redact_string(value)) + return _drop_encoded_credential(python_redact_string(value)) head: Final = value[:_MAX_SCRUBBED_ACCESS_ARG] kept: Final = head[: max(head.rfind("?"), head.rfind("&"))] if "?" in head else head - scrubbed: Final = _drop_encoded_credential(_python_redact_string(kept)) + scrubbed: Final = _drop_encoded_credential(python_redact_string(kept)) return f"{scrubbed}... ({len(value) - len(kept)} more chars truncated) ..." @@ -383,14 +388,14 @@ def _python_process_diagnostic( ) -> tuple[str, str | None, str | None, tuple[str, ...], bool]: def process_text(text: str) -> str: collapsed: Final = _collapse_base64_runs(text, base64_limit) if base64_limit > 0 else text - scrubbed: Final = _python_redact_string(collapsed) if redact else collapsed + scrubbed: Final = python_redact_string(collapsed) if redact else collapsed return _truncate_for_stdout_log(scrubbed, text_limit) if 0 < text_limit < len(scrubbed) else scrubbed processed_message: Final = process_text(message) processed_exception: Final = process_text(exception) if exception is not None else None - processed_stack: Final = _python_redact_string(stack) if redact and stack is not None else stack + processed_stack: Final = python_redact_string(stack) if redact and stack is not None else stack processed_leaves: Final = tuple( - _python_redact_structured_value(key, text) if redact else text for key, text in leaves + python_redact_structured_value(key, text) if redact else text for key, text in leaves ) changed: Final = ( processed_message != message @@ -472,7 +477,7 @@ def _process_record(record: logging.LogRecord, *, base64_limit: int, text_limit: record.stack_info = processed_stack # rebind-ok: the Filter interface mutates the record processed_values: Final = iter(processed_leaves[: len(extra_leaves)]) for key, original, prepared in extras: - replacement: Final = _sort_processed_sets(original, _replace_string_leaves(prepared, processed_values)) + replacement = _sort_processed_sets(original, _replace_string_leaves(prepared, processed_values)) if not _scrubbing_changed_nothing(replacement, original): setattr(record, key, replacement) raw_color_changed: Final = ( @@ -489,12 +494,12 @@ def _redact_json_record(value: object) -> object: leaves: Final = tuple(_string_leaves(None, prepared)) candidate: Final = diagnostics.run( lambda native: native.process_diagnostic("", None, None, leaves, (True, 0, 0))[3], - lambda: tuple(_python_redact_structured_value(key, text) for key, text in leaves), + lambda: tuple(python_redact_structured_value(key, text) for key, text in leaves), ) replacements: Final = ( candidate if len(candidate) == len(leaves) - else tuple(_python_redact_structured_value(key, text) for key, text in leaves) + else tuple(python_redact_structured_value(key, text) for key, text in leaves) ) return _sort_processed_sets(value, _replace_string_leaves(prepared, iter(replacements))) @@ -791,7 +796,7 @@ class CorrelationPlainFormatter(logging.Formatter): def format(self, record: logging.LogRecord) -> str: rendered: Final = super().format(record) - formatted: Final = rendered if _is_redacted(record) else _redact_string(rendered) + formatted: Final = rendered if _is_redacted(record) else redact_string(rendered) trace_id: Final = getattr(record, "trace_id", None) session_id: Final = getattr(record, "session_id", None) if not trace_id and not session_id: @@ -1050,7 +1055,7 @@ def _get_uvicorn_json_log_config(): return log_config -def _turn_on_json(): +def turn_on_json() -> None: """ Turn on JSON logging @@ -1064,12 +1069,18 @@ def _turn_on_json(): _setup_json_exception_handlers(JsonFormatter()) -def _turn_on_debug(): +_turn_on_json = turn_on_json + + +def turn_on_debug() -> None: verbose_logger.setLevel(level=logging.DEBUG) # set package log to debug verbose_router_logger.setLevel(level=logging.DEBUG) # set router logs to debug verbose_proxy_logger.setLevel(level=logging.DEBUG) # set proxy logs to debug +_turn_on_debug = turn_on_debug + + def _disable_debugging(): """Disable the package, router, and proxy verbose loggers.""" verbose_logger.disabled = True @@ -1093,8 +1104,11 @@ def print_verbose(print_statement): pass -def _is_debugging_on() -> bool: +def is_debugging_on() -> bool: """ Returns True if debugging is on """ return verbose_logger.isEnabledFor(logging.DEBUG) or set_verbose is True + + +_is_debugging_on = is_debugging_on diff --git a/litellm/_redis.py b/litellm/_redis.py index 791fa4ce783..499d713e5e4 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -26,7 +26,7 @@ from litellm._redis_credential_provider import ( AzureADCredentialProvider, ElastiCacheIAMCredentialProvider, GCPIAMCredentialProvider, - _generate_gcp_iam_access_token, + generate_gcp_iam_access_token, ) from litellm.constants import ( REDIS_CLUSTER_HEALTH_CHECK_INTERVAL, @@ -37,6 +37,8 @@ from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from ._logging import verbose_logger +_generate_gcp_iam_access_token = generate_gcp_iam_access_token + AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default" _AWS_IAM_KWARG_NAMES: Final = ( @@ -341,7 +343,7 @@ def create_gcp_iam_redis_connect_func( self._parser.on_connect(self) - auth_args: Final = (_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: diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index 7d90f944657..8b20b9b7d4b 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -38,7 +38,7 @@ class AzureCredential(Protocol): def get_token(self, *scopes: str) -> AzureAccessToken: ... -def _generate_gcp_iam_access_token(service_account: str) -> str: +def generate_gcp_iam_access_token(service_account: str) -> str: """ Generate GCP IAM access token for Redis authentication. @@ -65,6 +65,9 @@ def _generate_gcp_iam_access_token(service_account: str) -> str: return str(response.access_token) +_generate_gcp_iam_access_token = generate_gcp_iam_access_token + + def _get_cached_gcp_iam_token(service_account: str) -> str: """ Return a cached GCP IAM token, refreshing only when expired. @@ -93,7 +96,7 @@ def _get_cached_gcp_iam_token(service_account: str) -> str: if time.monotonic() < expiry: return token - token = _generate_gcp_iam_access_token(service_account) + token = generate_gcp_iam_access_token(service_account) _token_cache[service_account] = ( token, time.monotonic() + _GCP_IAM_TOKEN_TTL_SECONDS, diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 23ef2a99585..2b55e10a1d3 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -3,7 +3,7 @@ from collections.abc import Iterable, Iterator, Mapping from dataclasses import dataclass from dataclasses import replace as dataclasses_replace from enum import Enum -from typing import Any, Final, Literal +from typing import Any, Final, Literal, cast import litellm from litellm._logging import verbose_logger @@ -70,7 +70,7 @@ def batch_cost_is_final(batch: Batch) -> bool: async def calculate_batch_cost_and_usage( - file_content_dictionary: list[dict], + file_content_dictionary: list[dict[str, object]], custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"], model_name: str | None = None, model_info: ModelInfo | None = None, @@ -96,11 +96,11 @@ async def calculate_batch_cost_and_usage( ) -async def _handle_completed_batch( +async def handle_completed_batch( batch: Batch, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"], model_name: str | None = None, - litellm_params: dict | None = None, + litellm_params: dict[str, object] | None = None, model_info: ModelInfo | None = None, ) -> BatchCostUsageResult: """Fetch a completed batch's output file and aggregate its cost, usage, and @@ -162,6 +162,9 @@ async def _handle_completed_batch( ) +_handle_completed_batch = handle_completed_batch + + class _LineOutcome(Enum): """A batch output line that yielded no billable stats.""" @@ -183,7 +186,7 @@ class _BatchOutputLineStats: def _classify_output_line_stats( - entries: Iterable[dict], + entries: Iterable[dict[str, object]], custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"], model_name: str | None, model_info: ModelInfo | None, @@ -295,7 +298,7 @@ def _output_line_cost( def _aggregate_batch_cost_usage_models( - entries: Iterable[dict], + entries: Iterable[dict[str, object]], custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"], model_name: str | None = None, model_info: ModelInfo | None = None, @@ -347,7 +350,7 @@ def _aggregate_batch_cost_usage_models( def calculate_vertex_ai_batch_cost_and_usage( - vertex_ai_batch_responses: Iterable[dict], + vertex_ai_batch_responses: Iterable[dict[str, object]], model_name: str | None = None, model_info: ModelInfo | None = None, ) -> BatchCostUsageResult: @@ -436,7 +439,7 @@ def _provider_output_file_id(output_file_id: str) -> str: async def _fetch_batch_managed_file_content( file_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai", - litellm_params: dict | None = None, + litellm_params: Mapping[str, object] | None = None, ) -> bytes: """ Fetch a batch's output or error file and return its raw JSONL bytes. @@ -448,17 +451,18 @@ async def _fetch_batch_managed_file_content( Required for Azure and other providers that need authentication """ from litellm.files.main import afile_content + from litellm.files.types import FileContentCallOptions, FileContentRequestKwargs - # Build kwargs for afile_content with credentials from litellm_params - file_content_kwargs: Final = { - "file_id": _provider_output_file_id(file_id), + provider_output_file_id: Final = _provider_output_file_id(file_id) + credentials: Final = extract_file_access_credentials(litellm_params) + file_content_kwargs: Final[FileContentRequestKwargs] = { + "file_id": provider_output_file_id, "custom_llm_provider": custom_llm_provider, + **cast( # cast-ok: preserve dynamic provider credentials without validation + FileContentCallOptions, credentials + ), } - # Extract and add credentials for file access - credentials: Final = _extract_file_access_credentials(litellm_params) - file_content_kwargs.update(credentials) - _file_content: Final = await afile_content(**file_content_kwargs) return _file_content.content @@ -466,7 +470,7 @@ async def _fetch_batch_managed_file_content( async def _fetch_batch_output_file_content( batch: Batch, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai", - litellm_params: dict | None = None, + litellm_params: Mapping[str, object] | None = None, ) -> bytes: """ Fetch the batch output file and return its raw JSONL bytes @@ -488,7 +492,7 @@ async def _fetch_batch_output_file_content( async def count_error_file_failed_requests( batch: Batch, custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"], - litellm_params: dict | None, + litellm_params: Mapping[str, object] | None, ) -> int: """Count failed requests reported only in the batch's separate error file. @@ -506,10 +510,12 @@ async def count_error_file_failed_requests( except Exception as e: # noqa: BLE001 # a failed/missing error file must not abort cost tracking for the batch verbose_logger.debug("Failed to fetch batch error file %s: %s", batch.error_file_id, e) return 0 - return sum(1 for _ in _iter_batch_input_lines(error_file_content)) + return sum(1 for _ in iter_batch_input_lines(error_file_content)) -def _extract_file_access_credentials(litellm_params: dict | None) -> dict: +def extract_file_access_credentials( + litellm_params: Mapping[str, object] | None, +) -> dict[str, object]: """ Extract credentials from litellm_params for file access operations. @@ -522,7 +528,7 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict: Returns: Dictionary containing only the credentials needed for file access """ - credentials: Final = {} + credentials: Final[dict[str, object]] = {} if litellm_params: # List of credential keys that should be passed to file operations @@ -555,7 +561,10 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict: return credentials -def _get_file_content_as_dictionary(file_content: bytes) -> list[dict]: +_extract_file_access_credentials = extract_file_access_credentials + + +def get_file_content_as_dictionary(file_content: bytes) -> list[dict[str, object]]: """ Get the file content as a list of dictionaries from JSON Lines format, skipping malformed lines @@ -563,7 +572,10 @@ def _get_file_content_as_dictionary(file_content: bytes) -> list[dict]: return list(_iter_batch_output_entries(file_content)) -def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]: +_get_file_content_as_dictionary = get_file_content_as_dictionary + + +def iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]: """ Yield non-empty JSONL lines (unparsed) one at a time, so a caller can parse each row in its own try/except and a single malformed line cannot abort the @@ -581,27 +593,30 @@ def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]: yield line -def _iter_batch_output_entries(file_content: bytes) -> Iterator[dict]: +_iter_batch_input_lines = iter_batch_input_lines + + +def _iter_batch_output_entries(file_content: bytes) -> Iterator[dict[str, object]]: """ Yield parsed batch output JSONL entries one at a time without materializing the whole file as a list, so peak memory stays bounded. A malformed or non-object line is skipped with a warning so one bad line never aborts the whole batch's cost accounting. """ - for line in _iter_batch_input_lines(file_content): + for line in iter_batch_input_lines(file_content): entry = _parse_batch_output_line(line) if entry is not None: yield entry -def _parse_batch_output_line(line: bytes) -> dict | None: +def _parse_batch_output_line(line: bytes) -> dict[str, object] | None: try: parsed: Final[object] = json.loads(line) except ValueError as e: verbose_logger.warning("skipping malformed batch output line: %s", str(e)) return None if isinstance(parsed, dict): - return parsed + return cast("dict[str, object]", parsed) # cast-ok: JSON object keys are strings by definition verbose_logger.warning("skipping non-object batch output line of type %s", type(parsed).__name__) return None @@ -611,24 +626,29 @@ def _parse_batch_output_line(line: bytes) -> dict | None: _BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN: Final = 4 -def _estimate_batch_entry_tokens(raw_line: bytes) -> int: +def estimate_batch_entry_tokens(raw_line: bytes) -> int: """Conservative token estimate for a batch row the token counter cannot measure (or that cannot be parsed). Keeps the batch token total non-zero so a crafted row cannot evade the TPM limit, without hard-rejecting a legitimate batch.""" return max(1, len(raw_line) // _BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN) -def _count_entry_tokens( - entry: dict, +_estimate_batch_entry_tokens = estimate_batch_entry_tokens + + +def count_entry_tokens( + entry: Mapping[str, object], model_name: str | None = None, ) -> int: """Token-count a single batch input entry's body (chat / text / embedding).""" - body: Final = entry.get("body", {}) or {} - model: Final = body.get("model", model_name or "") + body: Final = cast( # cast-ok: batch payload bodies come from provider JSON + Mapping[str, object], entry.get("body", {}) or {} + ) + model: Final = cast(str, body.get("model", model_name or "")) # cast-ok: provider batch model names are strings messages: Final = body.get("messages") if messages: - return token_counter(model=model, messages=messages) + return token_counter(model=model, messages=cast(list[dict[str, object]], messages)) prompt: Final = body.get("prompt") if prompt: @@ -641,6 +661,9 @@ def _count_entry_tokens( return 0 +_count_entry_tokens = count_entry_tokens + + def _count_prompt_or_input_tokens(model: str, value: object) -> int: """Token-count a ``prompt`` / ``input`` field that the OpenAI batch schema allows in four shapes: @@ -707,8 +730,8 @@ def _get_batch_job_usage_from_response_body( from litellm.responses.utils import ResponseAPILoggingUtils _usage_dict: Final = response_body.get("usage", None) or {} - if ResponseAPILoggingUtils._is_response_api_usage(_usage_dict): - return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_usage_dict) + if ResponseAPILoggingUtils.is_response_api_usage(_usage_dict): + return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(_usage_dict) usage: Final[Usage] = Usage(**_usage_dict) if custom_llm_provider == "xai": from litellm.llms.xai.chat.transformation import XAIChatConfig @@ -737,9 +760,11 @@ def _get_response_from_batch_job_output_file( 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: Final[dict] = batch_job_output_file.get("response", None) or {} + _response: Final = cast( # cast-ok: batch output response comes from provider JSON + Mapping[str, object], batch_job_output_file.get("response", None) or {} + ) _response_body: Final = _response.get("body", None) or {} - return _response_body + return cast(Mapping[str, object], _response_body) # cast-ok: batch response body comes from provider JSON def _batch_response_was_successful( @@ -756,5 +781,7 @@ def _batch_response_was_successful( 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: Final[dict] = batch_job_output_file.get("response", None) or {} + _response: Final[dict[str, object]] = cast( # cast-ok: batch response bodies come from provider JSON + dict[str, object], batch_job_output_file.get("response", None) or {} + ) return _response.get("status_code", None) == 200 diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 888adffda7f..1dc4de04dc1 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -16,7 +16,7 @@ import traceback from collections.abc import Generator, Mapping from contextlib import contextmanager from enum import Enum -from typing import Any, Final, Literal +from typing import Any, Final, Literal, cast from pydantic import BaseModel @@ -363,12 +363,12 @@ class Cache: cache_key = "" # verbose_logger.debug("\nGetting Cache key. Kwargs: %s", kwargs) - preset_cache_key: Final = 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: Final = ModelParamHelper._get_all_llm_api_params() + combined_kwargs: Final = ModelParamHelper.get_all_llm_api_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: @@ -460,7 +460,7 @@ class Cache: or litellm_params.get("file_name") ) - def _get_preset_cache_key_from_kwargs(self, **kwargs) -> str | None: + def get_preset_cache_key_from_kwargs(self, **kwargs: object) -> str | None: """ Get the preset cache key from kwargs["litellm_params"] @@ -469,11 +469,17 @@ class Cache: 1. optional params like max_tokens, get transformed for bedrock -> max_new_tokens 2. avoid doing duplicate / repeated work """ - if kwargs: - if "litellm_params" in kwargs: - return kwargs["litellm_params"].get("preset_cache_key", None) + if "litellm_params" in kwargs: + litellm_params: Final = cast( # cast-ok: cache kwargs retain dynamic caller values + Mapping[str, object], kwargs["litellm_params"] + ) + return cast( # cast-ok: preserve dynamically supplied cache keys + str | None, litellm_params.get("preset_cache_key", None) + ) return None + _get_preset_cache_key_from_kwargs = get_preset_cache_key_from_kwargs + def _set_preset_cache_key_in_kwargs(self, preset_cache_key: str, **kwargs) -> None: """ Set the calculated cache key in kwargs @@ -539,11 +545,11 @@ class Cache: } time.sleep(CACHED_STREAMING_CHUNK_DELAY) - def _get_cache_logic( + def get_cache_logic( self, cached_result: object | None, max_age: float | None, - ): + ) -> object | None: """ Common get cache logic across sync + async implementations """ @@ -572,6 +578,8 @@ class Cache: return cached_response return cached_result + _get_cache_logic = get_cache_logic + @staticmethod def _get_safe_cache_lookup_kwargs(kwargs: Mapping[str, object]) -> dict[str, object]: cache_lookup_kwargs: Final[dict[str, object]] = {} @@ -628,7 +636,7 @@ class Cache: original_kwargs=kwargs, cache_lookup_kwargs=cache_lookup_kwargs, ) - return self._get_cache_logic(cached_result=cached_result, max_age=max_age) + return self.get_cache_logic(cached_result=cached_result, max_age=max_age) except Exception: print_verbose(f"An exception occurred: {traceback.format_exc()}") return None @@ -658,7 +666,7 @@ class Cache: cached_result = await dynamic_cache_object.async_get_cache(cache_key, **kwargs) else: cached_result = await self.cache.async_get_cache(cache_key, **kwargs) - return self._get_cache_logic(cached_result=cached_result, max_age=max_age) + return self.get_cache_logic(cached_result=cached_result, max_age=max_age) except Exception: print_verbose(f"An exception occurred: {traceback.format_exc()}") return None @@ -957,7 +965,7 @@ class Cache: if hasattr(self.cache, "disconnect"): await self.cache.disconnect() - def _supports_async(self) -> bool: + def supports_async(self) -> bool: """ Internal method to check if the cache type supports async get/set operations @@ -966,6 +974,8 @@ class Cache: """ return True + _supports_async = supports_async + def enable_cache( type: LiteLLMCacheType | None = LiteLLMCacheType.LOCAL, diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index a0bae6685de..4cb84eafa9e 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -33,7 +33,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( update_response_metadata, ) from litellm.litellm_core_utils.logging_utils import ( - _assemble_complete_response_from_streaming_chunks, + assemble_complete_response_from_streaming_chunks, ) from litellm.types.caching import CACHED_STREAM_EVENTS_KEY, EMBEDDING_CACHE_FORMAT_VERSION, CachedEmbedding from litellm.types.integrations.custom_logger import converted_stream_requested @@ -66,7 +66,7 @@ _StreamResultT = TypeVar("_StreamResultT") from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, + get_parent_otel_span_from_kwargs, ) from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper @@ -246,14 +246,14 @@ def _current_format_embedding_entry(entry: object) -> CachedEmbedding | None: class LLMCachingHandler: def __init__( self, - original_function: Callable, + original_function: Callable[..., object], request_kwargs: dict[str, object], start_time: datetime.datetime, ): from litellm.caching import DualCache, RedisCache - self.async_streaming_chunks: list[ModelResponse] = [] - self.sync_streaming_chunks: list[ModelResponse] = [] + self.async_streaming_chunks: list[object] = [] + self.sync_streaming_chunks: list[object] = [] self.request_kwargs = _drop_logging_obj_from_kwargs(request_kwargs) self.preset_cache_key: str | None = None self.original_function = original_function @@ -266,10 +266,10 @@ class LLMCachingHandler: else: self.dual_cache = None - async def _async_get_cache( + async def async_get_cache( self, model: str, - original_function: Callable, + original_function: Callable[..., object], logging_obj: LiteLLMLoggingObj, start_time: datetime.datetime, call_type: str, @@ -313,7 +313,7 @@ class LLMCachingHandler: cache_check_start_time: Final = time.perf_counter() cache_check_end_time: float | None = None ######################################################### - parent_otel_span: Final = _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): @@ -406,10 +406,12 @@ class LLMCachingHandler: # Caching disabled - return None to indicate no caching attempted return None - def _sync_get_cache( + _async_get_cache = async_get_cache + + def sync_get_cache( self, model: str, - original_function: Callable, + original_function: Callable[..., object], logging_obj: LiteLLMLoggingObj, start_time: datetime.datetime, call_type: str, @@ -496,6 +498,8 @@ class LLMCachingHandler: return CachingHandlerResponse(cached_result=cached_result) return CachingHandlerResponse(cached_result=cached_result) + _sync_get_cache = sync_get_cache + def handle_kwargs_input_list_or_str(self, kwargs: dict[str, object]) -> list[str]: """ Handles the input of kwargs['input'] being a list or a string @@ -715,7 +719,7 @@ class LLMCachingHandler: except Exception: return None - def _combine_cached_embedding_response_with_api_result( + def combine_cached_embedding_response_with_api_result( self, _caching_handler_response: CachingHandlerResponse, embedding_response: EmbeddingResponse, @@ -763,6 +767,8 @@ class LLMCachingHandler: merged._response_ms = (end_time - start_time).total_seconds() * 1000 return merged + _combine_cached_embedding_response_with_api_result = combine_cached_embedding_response_with_api_result + def _async_log_cache_hit_on_callbacks( self, logging_obj: LiteLLMLoggingObj, @@ -861,7 +867,7 @@ class LLMCachingHandler: request_kwargs: Final = new_kwargs.copy() request_cache_key: Final = _request_cache_key(request_kwargs) request_kwargs.pop("cache_key", None) - if litellm.cache._supports_async() is True: + 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) with response_cache_phase("get"): @@ -1111,7 +1117,7 @@ class LLMCachingHandler: None """ from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, + get_parent_otel_span_from_kwargs, ) if litellm.cache is None: @@ -1128,10 +1134,10 @@ class LLMCachingHandler: args, ) ) - parent_otel_span: Final = _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): + if self.should_store_result_in_cache(original_function=original_function, kwargs=new_kwargs): if ( isinstance(result, litellm.ModelResponse) or isinstance(result, litellm.EmbeddingResponse) @@ -1183,13 +1189,13 @@ class LLMCachingHandler: ) ) - if self._should_store_result_in_cache(original_function=self.original_function, kwargs=new_kwargs): + if self.should_store_result_in_cache(original_function=self.original_function, kwargs=new_kwargs): with response_cache_phase("set"): litellm.cache.add_cache(result, **new_kwargs) return - def _should_store_result_in_cache(self, original_function: Callable, kwargs: dict[str, Any]) -> bool: + def should_store_result_in_cache(self, original_function: Callable[..., object], kwargs: dict[str, Any]) -> bool: """ Helper function to determine if the result should be stored in the cache. @@ -1200,6 +1206,8 @@ class LLMCachingHandler: kwargs.get("cache", {}).get("no-store", False) is not True ) + _should_store_result_in_cache = should_store_result_in_cache + def wrap_streaming_result_for_cache( self, result: _StreamResultT, call_type: str ) -> "_StreamResultT | AnthropicMessagesStreamCacheWriter": @@ -1208,7 +1216,7 @@ class LLMCachingHandler: CallTypes.aanthropic_messages.value, ): return result - if litellm.cache is None or not self._should_store_result_in_cache( + if litellm.cache is None or not self.should_store_result_in_cache( original_function=self.original_function, kwargs=self.request_kwargs ): return result @@ -1240,7 +1248,7 @@ class LLMCachingHandler: covering_call_types: Final = ("aresponses", "responses") if call_type == "aresponses" else (call_type,) return any(name in litellm.cache.supported_call_types for name in covering_call_types) - async def _add_streaming_response_to_cache(self, processed_chunk: ModelResponse): + async def add_streaming_response_to_cache(self, processed_chunk: ModelResponse) -> None: """ Internal method to add the streaming response to the cache @@ -1251,7 +1259,7 @@ class LLMCachingHandler: """ complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = ( - _assemble_complete_response_from_streaming_chunks( + assemble_complete_response_from_streaming_chunks( result=processed_chunk, start_time=self.start_time, end_time=datetime.datetime.now(), @@ -1268,12 +1276,14 @@ class LLMCachingHandler: kwargs=self.request_kwargs, ) - def _sync_add_streaming_response_to_cache(self, processed_chunk: ModelResponse): + _add_streaming_response_to_cache = add_streaming_response_to_cache + + def sync_add_streaming_response_to_cache(self, processed_chunk: ModelResponse) -> None: """ Sync internal method to add the streaming response to the cache """ complete_streaming_response: Final[ModelResponse | TextCompletionResponse | None] = ( - _assemble_complete_response_from_streaming_chunks( + assemble_complete_response_from_streaming_chunks( result=processed_chunk, start_time=self.start_time, end_time=datetime.datetime.now(), @@ -1290,6 +1300,8 @@ class LLMCachingHandler: kwargs=self.request_kwargs, ) + _sync_add_streaming_response_to_cache = sync_add_streaming_response_to_cache + def _update_litellm_logging_obj_environment( self, logging_obj: LiteLLMLoggingObj, @@ -1329,7 +1341,7 @@ class LLMCachingHandler: if litellm.cache is not None: litellm_params["preset_cache_key"] = ( - self.preset_cache_key or litellm.cache._get_preset_cache_key_from_kwargs(**kwargs) + self.preset_cache_key or litellm.cache.get_preset_cache_key_from_kwargs(**kwargs) ) else: litellm_params["preset_cache_key"] = None diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 2ee4bff6112..b7828ac7871 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -38,7 +38,7 @@ from litellm.constants import ( REDIS_CIRCUIT_BREAKER_TIMEOUT_MIN_DURATION, REDIS_TIMEOUT_LOG_INTERVAL, ) -from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs +from litellm.litellm_core_utils.core_helpers import get_parent_otel_span_from_kwargs from litellm.litellm_core_utils.coroutine_checker import coroutine_checker from litellm.types.caching import ( RedisPipelineIncrementOperation, @@ -1157,7 +1157,7 @@ class RedisCache(BaseCache): error=e, start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), call_type="async_set_cache", caller=_get_call_stack_info(), ) @@ -1192,7 +1192,7 @@ class RedisCache(BaseCache): caller=_get_call_stack_info(), start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) return result @@ -1208,7 +1208,7 @@ class RedisCache(BaseCache): caller=_get_call_stack_info(), start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) log_redis_failure( @@ -1276,7 +1276,7 @@ class RedisCache(BaseCache): caller=_get_call_stack_info(), start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) return @@ -1293,7 +1293,7 @@ class RedisCache(BaseCache): caller=_get_call_stack_info(), start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) @@ -1384,7 +1384,7 @@ class RedisCache(BaseCache): error=e, start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), call_type="async_set_cache_sadd", caller=_get_call_stack_info(), ) @@ -1410,7 +1410,7 @@ class RedisCache(BaseCache): caller=_get_call_stack_info(), start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) except Exception as e: @@ -1425,7 +1425,7 @@ class RedisCache(BaseCache): caller=_get_call_stack_info(), start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) # NON blocking - notify users Redis is throwing an exception @@ -2036,7 +2036,7 @@ class RedisCache(BaseCache): caller=_get_call_stack_info(), start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) return results @@ -2053,7 +2053,7 @@ class RedisCache(BaseCache): caller=_get_call_stack_info(), start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) log_redis_failure( diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 93d79bb3ac8..6af4aef38b0 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -244,7 +244,7 @@ def tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _Cha 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")), + id=LiteLLMCompletionResponsesConfig.tool_call_id_from_responses_item(item.get("id"), item.get("call_id")), type="function", function=function_chunk, index=index, @@ -525,7 +525,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): elif key == "previous_response_id": responses_api_request["previous_response_id"] = value elif key == "reasoning_effort": - responses_api_request["reasoning"] = self._map_reasoning_effort(value) + responses_api_request["reasoning"] = self.map_reasoning_effort(value) elif key == "web_search_options": self._add_web_search_tool(responses_api_request, value) @@ -983,7 +983,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): setattr( model_response, "usage", - ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_response.usage), + ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(raw_response.usage), ) model_response.id = _upstream_response_id(raw_response.id) or raw_response.id @@ -1206,7 +1206,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return optional_params - def _map_reasoning_effort(self, reasoning_effort: object) -> Reasoning: + def map_reasoning_effort(self, reasoning_effort: object) -> Reasoning: # If dict is passed, convert it directly to Reasoning object if isinstance(reasoning_effort, dict): return Reasoning( @@ -1225,6 +1225,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): else Reasoning(effort=reasoning_effort) ) + _map_reasoning_effort = map_reasoning_effort + def _add_web_search_tool( self, responses_api_request: ResponsesAPIOptionalRequestParams, @@ -1626,7 +1628,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): if response_data.get("usage"): from litellm.responses.utils import ResponseAPILoggingUtils - usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response_data.get("usage")) + usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(response_data.get("usage")) provider_metadata: Final = _provider_metadata(response_data) served_service_tier: Final = response_data.get("service_tier") return ModelResponseStream( diff --git a/litellm/constants.py b/litellm/constants.py index 95d05c6d819..6343c0675e2 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1077,7 +1077,7 @@ openai_text_completion_compatible_providers: Final[list] = [ # providers that s "hyperbolic", "wandb", ] -_openai_like_providers: Final[list] = [ +_openai_like_providers: Final[list[str]] = [ "predibase", "databricks", "lemonade", diff --git a/litellm/containers/utils.py b/litellm/containers/utils.py index ed604a53e1c..44923719543 100644 --- a/litellm/containers/utils.py +++ b/litellm/containers/utils.py @@ -19,7 +19,7 @@ def decode_managed_container_id_for_request( Returns: (original_container_id, resolved_provider, updated_litellm_params) """ - decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) + decoded: Final = ResponsesAPIRequestUtils.decode_container_id(container_id) original_container_id: Final = decoded.get("response_id", container_id) decoded_provider: Final = decoded.get("custom_llm_provider") @@ -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: Final = 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 5a4b5bcb339..10d61087ed9 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -29,13 +29,13 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import ( _SERVICE_TIER_TO_COST_KEY_SUFFIX, BilledTokenRates, CostCalculatorUtils, - _generic_cost_per_character, - _get_regional_uplift_multiplier, - _get_service_tier_cost_key, calculate_cost_component, + generic_cost_per_character, generic_cost_per_token, get_batch_cost_rates, get_billable_input_tokens, + get_regional_uplift_multiplier, + get_service_tier_cost_key, get_token_type_cost_breakdown, parse_prompt_tokens_details, select_cost_metric_for_model, @@ -137,7 +137,7 @@ from litellm.utils import ( ProviderConfigManager, TextCompletionResponse, TranscriptionResponse, - _cached_get_model_info_helper, + cached_get_model_info_helper, token_counter, ) @@ -349,7 +349,7 @@ def _per_second_pricing_cost( audio_seconds: float = 0.0, ) -> tuple[float, float] | None: try: - model_info: Final = _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) except Exception: # noqa: BLE001 # the lookup raises plain Exception for an unmapped model return None if _has_token_or_tiered_pricing(model_info) or not _bills_wall_clock_seconds(model_info): @@ -587,7 +587,7 @@ def cost_per_token( raise ValueError( f"prompt_characters must be provided for tts calls. prompt_characters={prompt_characters}, model={model}, custom_llm_provider={custom_llm_provider}, call_type={call_type}" ) - _prompt_cost, _completion_cost = _generic_cost_per_character( + _prompt_cost, _completion_cost = generic_cost_per_character( model=model_without_prefix, custom_llm_provider=custom_llm_provider, prompt_characters=prompt_characters, @@ -749,7 +749,7 @@ def cost_per_token( service_tier=service_tier, ) else: - model_info: Final = _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 _has_token_or_tiered_pricing(model_info): return generic_cost_per_token( model=model, @@ -818,7 +818,7 @@ def _cost_map_entry_prices_anything(entry: Mapping[str, object]) -> bool: ) -def _select_model_name_for_cost_calc( +def select_model_name_for_cost_calc( model: str | None, completion_response: object | None, base_model: str | None = None, @@ -889,6 +889,9 @@ def _select_model_name_for_cost_calc( return return_model +_select_model_name_for_cost_calc = select_model_name_for_cost_calc + + def _strip_unregistered_leading_segments(model: str, region_name: str | None) -> str: """Resolve a provider-prefixed slash alias like "vertex_ai/vertex/claude-opus-5" to the registered cost key ("vertex_ai/claude-opus-5"), keeping the model unchanged when it already @@ -1039,9 +1042,9 @@ def get_usage_object( elif ( usage_obj is not None and (isinstance(usage_obj, dict) or isinstance(usage_obj, ResponseAPIUsage)) - and ResponseAPILoggingUtils._is_response_api_usage(usage_obj) + and ResponseAPILoggingUtils.is_response_api_usage(usage_obj) ): - return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage_obj) + return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage_obj) elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(usage_obj): return TranscriptionUsageObjectTransformation.transform_transcription_usage_object( cast( @@ -1070,7 +1073,7 @@ def _is_known_usage_objects(usage_obj): ) -def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: object) -> CallTypesLiteral | None: +def infer_call_type(call_type: CallTypesLiteral | None, completion_response: object) -> CallTypesLiteral | None: if call_type is not None: return call_type @@ -1097,6 +1100,9 @@ def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: ob return call_type +_infer_call_type = infer_call_type + + def _apply_cost_discount( base_cost: float, custom_llm_provider: str | None, @@ -1369,7 +1375,7 @@ def completion_cost( - For un-mapped Replicate models, the cost is calculated based on the total time used for the request. """ try: - call_type = _infer_call_type(call_type, completion_response) or "completion" + call_type = infer_call_type(call_type, completion_response) or "completion" if call_type == CallTypes.aresponses_websocket.value and isinstance( completion_response, LiteLLMRealtimeStreamLoggingObject @@ -1438,7 +1444,7 @@ def completion_cost( ) explicit_pricing: Final = custom_pricing is True or base_model is not None - selected_model: Final = _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, @@ -1492,10 +1498,8 @@ def completion_cost( .calculate_usage(usage_object=_usage, reasoning_content=None) .model_dump() ) - elif ResponseAPILoggingUtils._is_response_api_usage(_usage): - _usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( - _usage - ).model_dump() + elif ResponseAPILoggingUtils.is_response_api_usage(_usage): + _usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(_usage).model_dump() elif TranscriptionUsageObjectTransformation.is_transcription_usage_object(_usage): tr_usage = TranscriptionUsageObjectTransformation.transform_transcription_usage_object( cast( @@ -1561,7 +1565,7 @@ def completion_cost( "litellm.cost_calculator.py::completion_cost() - Error inferring custom_llm_provider - %s", e, ) - if CostCalculatorUtils._call_type_has_image_response(call_type) and isinstance( + if CostCalculatorUtils.call_type_has_image_response(call_type) and isinstance( completion_response, ImageResponse ): ### IMAGE GENERATION COST CALCULATION ### @@ -1631,7 +1635,7 @@ def completion_cost( video_resolution=video_resolution, ) elif call_type in _SPEECH_CALL_TYPES: - prompt_characters = litellm.utils._count_characters(text=prompt) + prompt_characters = litellm.utils.count_characters(text=prompt) elif call_type in _TRANSCRIPTION_CALL_TYPES: # Check _hidden_params first (duration stored there to # avoid polluting the response body), then fall back to @@ -1780,10 +1784,10 @@ def completion_cost( data={"messages": messages}, call_type="completion" ) - prompt_characters = litellm.utils._count_characters(text=prompt_string) + prompt_characters = litellm.utils.count_characters(text=prompt_string) if completion_response is not None and isinstance(completion_response, ModelResponse): completion_string = litellm.utils.get_response_string(response_obj=completion_response) - completion_characters = litellm.utils._count_characters(text=completion_string) + completion_characters = litellm.utils.count_characters(text=completion_string) # Get the original request model for router detection request_model_for_cost = None @@ -1946,7 +1950,7 @@ def get_response_cost_from_hidden_params( hidden_params: dict | BaseModel, ) -> float | None: if isinstance(hidden_params, BaseModel): - _hidden_params_dict = cast(BaseModel, hidden_params).model_dump() + _hidden_params_dict = hidden_params.model_dump() else: _hidden_params_dict = hidden_params @@ -2132,7 +2136,7 @@ def pricing_entry_for_cost_calc( deployment_key: Final = router_model_id or model if deployment_entry is not None and deployment_key is not None: return deployment_key, deployment_entry - selected_model: Final = _select_model_name_for_cost_calc( + selected_model: Final = select_model_name_for_cost_calc( model=model, completion_response=completion_response, base_model=base_model, @@ -2657,7 +2661,7 @@ def batch_cost_calculator( ) # batch cost is usually half of the regular token cost # Add cache read cost if applicable - cache_read_cost_key: Final = _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: Final = model_info.get("cache_creation_input_token_cost") or input_cost_per_token @@ -2672,7 +2676,7 @@ def batch_cost_calculator( text_tokens: Final = usage.completion_tokens - image_tokens total_completion_cost = text_tokens * text_rate + image_tokens * image_rate - uplift: Final = _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 @@ -2827,7 +2831,7 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor): ) usage_objects: Final[list[Usage]] = [] for result in response_done_events: - usage_object = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( + usage_object = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage( result["response"].get("usage", {}) ) usage_objects.append(usage_object) @@ -2885,9 +2889,7 @@ class ResponsesWebSocketTokenUsageProcessor(BaseTokenUsageProcessor): results: Sequence[Mapping[str, object]], ) -> tuple[Usage, ...]: return tuple( - ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( # pyright: ignore[reportPrivateUsage] # same shared transform the realtime processor uses - response.usage - ) + ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(response.usage) for _, response in _billable_responses_ws_events(results) if response.usage is not None ) diff --git a/litellm/files/main.py b/litellm/files/main.py index 9057211256d..94fdb41042b 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -30,9 +30,6 @@ FileCreateProvider = Literal[ "mistral", "xai", ] -FileRetrieveProvider = Literal[ - "openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral", "xai" -] FileDeleteProvider = Literal[ "openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic", "mistral", "xai" ] @@ -40,7 +37,7 @@ FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthrop import litellm from litellm import get_secret_str from litellm.files.streaming import FileContentStreamingResponse -from litellm.files.types import FileContentProvider, FileContentStreamingResult +from litellm.files.types import FileContentProvider, FileContentStreamingResult, FileRetrieveProvider from litellm.litellm_core_utils.get_litellm_params import add_trusted_model_credentials_to_litellm_params from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj diff --git a/litellm/files/types.py b/litellm/files/types.py index ae29ce2721f..9ce63ac9190 100644 --- a/litellm/files/types.py +++ b/litellm/files/types.py @@ -1,9 +1,37 @@ from collections.abc import AsyncIterator, Iterator, Mapping -from typing import Literal, NamedTuple +from typing import Literal, NamedTuple, TypedDict + +from typing_extensions import NotRequired, ReadOnly FileContentProvider = Literal[ "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "manus", "mistral" ] +FileRetrieveProvider = Literal[ + "openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral", "xai" +] + + +class FileContentCallOptions(TypedDict, total=False): + custom_llm_provider: ReadOnly[FileContentProvider] + extra_body: ReadOnly[NotRequired[dict[str, str] | None]] + extra_headers: ReadOnly[NotRequired[dict[str, str] | None]] + chunk_size: ReadOnly[NotRequired[int]] + stream: ReadOnly[NotRequired[bool]] + + +class FileContentRequestKwargs(TypedDict): + file_id: ReadOnly[str] + custom_llm_provider: ReadOnly[FileContentProvider] + extra_body: ReadOnly[NotRequired[dict[str, str] | None]] + extra_headers: ReadOnly[NotRequired[dict[str, str] | None]] + chunk_size: ReadOnly[NotRequired[int]] + stream: ReadOnly[NotRequired[bool]] + + +class FileRetrieveCallOptions(TypedDict, total=False): + custom_llm_provider: ReadOnly[FileRetrieveProvider] + extra_body: ReadOnly[NotRequired[dict[str, str] | None]] + extra_headers: ReadOnly[NotRequired[dict[str, str] | None]] class FileContentStreamingResult(NamedTuple): diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index cac8c7cb337..704a25b0457 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -31,7 +31,7 @@ from litellm.integrations.SlackAlerting.hanging_request_check import ( ) from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.litellm_core_utils.exception_mapping_utils import ( - _add_key_name_and_team_to_alert, + add_key_name_and_team_to_alert, ) from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -293,7 +293,7 @@ class SlackAlerting(CustomBatchLogger): # add deployment latencies to alert if kwargs is not None and "litellm_params" in kwargs and "metadata" in kwargs["litellm_params"]: _metadata: Final[dict] = kwargs["litellm_params"]["metadata"] - request_info = _add_key_name_and_team_to_alert(request_info=request_info, metadata=_metadata) + request_info = add_key_name_and_team_to_alert(request_info=request_info, metadata=_metadata) _deployment_latency_map: Final = self._get_deployment_latencies_to_alert(metadata=_metadata) if _deployment_latency_map is not None: diff --git a/litellm/integrations/SlackAlerting/utils.py b/litellm/integrations/SlackAlerting/utils.py index da256e5143f..a88fe3fde53 100644 --- a/litellm/integrations/SlackAlerting/utils.py +++ b/litellm/integrations/SlackAlerting/utils.py @@ -4,7 +4,7 @@ Utils used for slack alerting import asyncio from collections.abc import Callable -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Any, Final, cast import litellm from litellm.integrations.custom_logger import CustomLogger @@ -57,8 +57,8 @@ def process_slack_alerting_variables( return alert_to_webhook_url -async def _add_langfuse_trace_id_to_alert( - request_data: dict | None = None, +async def add_langfuse_trace_id_to_alert( + request_data: dict[str, object] | None = None, ) -> str | None: """ Returns langfuse trace url @@ -71,7 +71,7 @@ async def _add_langfuse_trace_id_to_alert( from litellm.integrations.langfuse.langfuse import LangFuseLogger, resolve_langfuse_host callbacks: Final[list[CustomLogger | Callable[..., object] | str]] = ( - litellm.logging_callback_manager._get_all_callbacks() + litellm.logging_callback_manager.get_all_callbacks() ) if not any(callback == "langfuse" or isinstance(callback, LangFuseLogger) for callback in callbacks): return None @@ -79,7 +79,9 @@ async def _add_langfuse_trace_id_to_alert( if request_data is None or request_data.get("litellm_logging_obj", None) is None: return None - litellm_logging_obj: Final[Logging] = request_data["litellm_logging_obj"] + litellm_logging_obj: Final = cast( # cast-ok: logging object crosses the dynamic callback payload boundary + Logging, request_data["litellm_logging_obj"] + ) instance_host: Final = next( (callback.langfuse_host for callback in callbacks if isinstance(callback, LangFuseLogger)), None ) @@ -87,8 +89,11 @@ async def _add_langfuse_trace_id_to_alert( litellm_logging_obj.standard_callback_dynamic_params.get("langfuse_host") or instance_host ) for _ in range(3): - if (trace_id := litellm_logging_obj._get_trace_id(service_name="langfuse")) is not None: + if (trace_id := litellm_logging_obj.get_trace_id(service_name="langfuse")) is not None: return f"{host}/trace/{trace_id}" await asyncio.sleep(3) # wait 3s before retrying for trace id return None + + +_add_langfuse_trace_id_to_alert = add_langfuse_trace_id_to_alert diff --git a/litellm/integrations/arize/arize_phoenix_prompt_manager.py b/litellm/integrations/arize/arize_phoenix_prompt_manager.py index fbf25ea87fd..3bc33bb68a2 100644 --- a/litellm/integrations/arize/arize_phoenix_prompt_manager.py +++ b/litellm/integrations/arize/arize_phoenix_prompt_manager.py @@ -118,9 +118,9 @@ class ArizePhoenixTemplateManager: # Load prompt from Arize Phoenix if prompt_id is provided if self.prompt_id: - self._load_prompt_from_arize(self.prompt_id) + self.load_prompt_from_arize(self.prompt_id) - def _load_prompt_from_arize(self, prompt_version_id: str) -> None: + def load_prompt_from_arize(self, prompt_version_id: str) -> None: """Load a specific prompt version from Arize Phoenix.""" try: # Fetch the prompt version from Arize Phoenix @@ -134,6 +134,8 @@ class ArizePhoenixTemplateManager: except Exception as e: raise Exception(f"Failed to load prompt version '{prompt_version_id}' from Arize Phoenix: {e}") + _load_prompt_from_arize = load_prompt_from_arize + 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: Final[ArizePhoenixTemplateBody] = data.get("template", {}) @@ -418,7 +420,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement): try: # Load the prompt from Arize Phoenix if not already loaded if prompt_id not in self.prompt_manager.prompts: - self.prompt_manager._load_prompt_from_arize(prompt_id) + self.prompt_manager.load_prompt_from_arize(prompt_id) # Get the rendered messages and metadata rendered_messages, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables) diff --git a/litellm/integrations/bitbucket/bitbucket_prompt_manager.py b/litellm/integrations/bitbucket/bitbucket_prompt_manager.py index e98fa77a562..fc885b7ee2c 100644 --- a/litellm/integrations/bitbucket/bitbucket_prompt_manager.py +++ b/litellm/integrations/bitbucket/bitbucket_prompt_manager.py @@ -94,9 +94,9 @@ class BitBucketTemplateManager: # Load prompts from BitBucket if prompt_id is provided if self.prompt_id: - self._load_prompt_from_bitbucket(self.prompt_id) + self.load_prompt_from_bitbucket(self.prompt_id) - def _load_prompt_from_bitbucket(self, prompt_id: str) -> None: + def load_prompt_from_bitbucket(self, prompt_id: str) -> None: """Load a specific .prompt file from BitBucket.""" try: # Fetch the .prompt file from BitBucket @@ -108,6 +108,8 @@ class BitBucketTemplateManager: except Exception as e: raise Exception(f"Failed to load prompt '{prompt_id}' from BitBucket: {e}") + _load_prompt_from_bitbucket = load_prompt_from_bitbucket + def _parse_prompt_file(self, content: str, prompt_id: str) -> BitBucketPromptTemplate: """Parse a .prompt file content and extract metadata and template.""" # Split frontmatter and content @@ -446,7 +448,7 @@ class BitBucketPromptManager(CustomPromptManagement): try: # Load the prompt from BitBucket if not already loaded if prompt_id not in self.prompt_manager.prompts: - self.prompt_manager._load_prompt_from_bitbucket(prompt_id) + self.prompt_manager.load_prompt_from_bitbucket(prompt_id) # Get the rendered prompt and metadata rendered_prompt, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables) diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index d4162369a35..6a4987298c4 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -991,7 +991,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac import litellm from litellm._logging import verbose_logger - all_callbacks: Final = 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"): diff --git a/litellm/integrations/dotprompt/__init__.py b/litellm/integrations/dotprompt/__init__.py index 1188bce27da..4cf6c0de2df 100644 --- a/litellm/integrations/dotprompt/__init__.py +++ b/litellm/integrations/dotprompt/__init__.py @@ -27,7 +27,7 @@ def set_global_prompt_directory(directory: str) -> None: litellm.global_prompt_directory = directory -def _get_prompt_data_from_dotprompt_content(dotprompt_content: str) -> dict: +def get_prompt_data_from_dotprompt_content(dotprompt_content: str) -> dict[str, object]: """ Get the prompt data from the dotprompt content. @@ -37,12 +37,15 @@ def _get_prompt_data_from_dotprompt_content(dotprompt_content: str) -> dict: # Parse the dotprompt content to extract frontmatter and content temp_manager: Final = PromptManager() - metadata, content = temp_manager._parse_frontmatter(dotprompt_content) + metadata, content = temp_manager.parse_frontmatter(dotprompt_content) # Convert to prompt_data format return {"content": content.strip(), "metadata": metadata} +_get_prompt_data_from_dotprompt_content = get_prompt_data_from_dotprompt_content + + def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec") -> "CustomPromptManagement": """ Initialize a prompt from a .prompt file. @@ -60,7 +63,7 @@ def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "Prom # Handle dotprompt_content from database 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) + prompt_data = get_prompt_data_from_dotprompt_content(dotprompt_content) from .prompt_manager import strip_version_suffix diff --git a/litellm/integrations/dotprompt/prompt_manager.py b/litellm/integrations/dotprompt/prompt_manager.py index 9c82ff7c5ba..47c19b12b65 100644 --- a/litellm/integrations/dotprompt/prompt_manager.py +++ b/litellm/integrations/dotprompt/prompt_manager.py @@ -168,7 +168,7 @@ class PromptManager: content: Final = file_path.read_text(encoding="utf-8") # Split frontmatter and content - frontmatter, template_content = self._parse_frontmatter(content) + frontmatter, template_content = self.parse_frontmatter(content) return PromptTemplate( content=template_content.strip(), @@ -176,7 +176,7 @@ class PromptManager: template_id=prompt_id, ) - def _parse_frontmatter(self, content: str) -> tuple[dict[str, object], str]: + def parse_frontmatter(self, content: str) -> tuple[dict[str, object], str]: """Parse YAML frontmatter from prompt content.""" # Match YAML frontmatter between --- delimiters frontmatter_pattern: Final = r"^---\s*\n(.*?)\n---\s*\n(.*)$" @@ -197,6 +197,8 @@ class PromptManager: return frontmatter, template_content + _parse_frontmatter = parse_frontmatter + def render( self, prompt_id: str, @@ -329,7 +331,7 @@ class PromptManager: content: Final = file_path.read_text(encoding="utf-8") # Parse frontmatter and content - frontmatter, template_content = self._parse_frontmatter(content) + frontmatter, template_content = self.parse_frontmatter(content) return {"content": template_content.strip(), "metadata": frontmatter} diff --git a/litellm/integrations/dynamodb.py b/litellm/integrations/dynamodb.py index 3fbbfe91ddf..6569e70329b 100644 --- a/litellm/integrations/dynamodb.py +++ b/litellm/integrations/dynamodb.py @@ -3,7 +3,8 @@ import os import traceback -from collections.abc import Mapping +from collections.abc import Callable, Mapping +from datetime import datetime from typing import Final, Protocol import litellm @@ -32,9 +33,18 @@ class DyanmoDBLogger: ) self.table_name = litellm.dynamodb_table_name - async def _async_log_event(self, kwargs, response_obj, start_time, end_time, print_verbose): + async def async_log_event( + self, + kwargs: Mapping[str, object], + response_obj: object, + start_time: datetime, + end_time: datetime, + print_verbose: Callable[[str], object], + ) -> None: self.log_event(kwargs, response_obj, start_time, end_time, print_verbose) + _async_log_event = async_log_event + def log_event(self, kwargs, response_obj, start_time, end_time, print_verbose): try: print_verbose(f"DynamoDB Logging - Enters logging function for model {kwargs}") diff --git a/litellm/integrations/focus/export_engine.py b/litellm/integrations/focus/export_engine.py index 16f11a9e015..615aba72954 100644 --- a/litellm/integrations/focus/export_engine.py +++ b/litellm/integrations/focus/export_engine.py @@ -2,14 +2,15 @@ from __future__ import annotations -from typing import Any, Final +from collections.abc import Callable +from typing import Any, Final, SupportsFloat, SupportsIndex, SupportsInt, cast import polars as pl from litellm._logging import verbose_logger from .database import FocusLiteLLMDatabase -from .destinations import FocusDestinationFactory, FocusTimeWindow +from .destinations import FocusDestination, FocusDestinationFactory, FocusTimeWindow from .serializers import FocusCsvSerializer, FocusParquetSerializer, FocusSerializer from .transformer import FocusTransformer @@ -28,14 +29,46 @@ class FocusExportEngine: self.provider = provider self.export_format = export_format self.prefix = prefix - self._destination = FocusDestinationFactory.create( + self.destination = FocusDestinationFactory.create( provider=self.provider, prefix=self.prefix, config=destination_config, ) - self._serializer = self._init_serializer() - self._transformer = FocusTransformer() - self._database = FocusLiteLLMDatabase() + self.serializer = self._init_serializer() + self.transformer = FocusTransformer() + self.database = FocusLiteLLMDatabase() + + @property + def _destination(self) -> FocusDestination: + return self.destination + + @_destination.setter + def _destination(self, value: FocusDestination) -> None: + self.destination = value + + @property + def _serializer(self) -> FocusSerializer: + return self.serializer + + @_serializer.setter + def _serializer(self, value: FocusSerializer) -> None: + self.serializer = value + + @property + def _transformer(self) -> FocusTransformer: + return self.transformer + + @_transformer.setter + def _transformer(self, value: FocusTransformer) -> None: + self.transformer = value + + @property + def _database(self) -> FocusLiteLLMDatabase: + return self.database + + @_database.setter + def _database(self, value: FocusLiteLLMDatabase) -> None: + self.database = value def _init_serializer(self) -> FocusSerializer: if self.export_format == "csv": @@ -45,18 +78,18 @@ class FocusExportEngine: raise NotImplementedError(f"Export format '{self.export_format}' not supported. Use 'parquet' or 'csv'.") async def dry_run_export_usage_data(self, limit: int | None) -> dict[str, Any]: - data: Final = await self._database.get_usage_data(limit=limit) - normalized: Final = self._transformer.transform(data) + data: Final = await self.database.get_usage_data(limit=limit) + normalized: Final = self.transformer.transform(data) usage_sample: Final = data.head(min(50, len(data))).to_dicts() normalized_sample: Final = normalized.head(min(50, len(normalized))).to_dicts() summary: Final = { "total_records": len(normalized), - "total_spend": self._sum_column(data, "spend"), - "total_tokens": self._sum_column(data, "total_tokens"), - "unique_teams": self._count_unique(data, "team_id"), - "unique_models": self._count_unique(data, "model"), + "total_spend": self.sum_column(data, "spend"), + "total_tokens": self.sum_column(data, "total_tokens"), + "unique_teams": self.count_unique(data, "team_id"), + "unique_models": self.count_unique(data, "model"), } return { @@ -71,12 +104,12 @@ class FocusExportEngine: limit: int | None, ) -> None: """Export all available data without time-window filtering.""" - data: Final = await self._database.get_usage_data(limit=limit) + data: Final = await self.database.get_usage_data(limit=limit) if data.is_empty(): verbose_logger.debug("Focus export: no usage data available") return - normalized: Final = self._transformer.transform(data) + normalized: Final = self.transformer.transform(data) if normalized.is_empty(): verbose_logger.debug("Focus export: normalized data empty") return @@ -98,7 +131,7 @@ class FocusExportEngine: window: FocusTimeWindow, limit: int | None, ) -> None: - data: Final = await self._database.get_usage_data( + data: Final = await self.database.get_usage_data( limit=limit, start_time_utc=window.start_time, end_time_utc=window.end_time, @@ -107,7 +140,7 @@ class FocusExportEngine: verbose_logger.debug("Focus export: no usage data for window %s", window) return - normalized: Final = self._transformer.transform(data) + normalized: Final = self.transformer.transform(data) if normalized.is_empty(): verbose_logger.debug("Focus export: normalized data empty for window %s", window) return @@ -115,39 +148,45 @@ class FocusExportEngine: await self._serialize_and_upload(normalized, window) async def _serialize_and_upload(self, frame: pl.DataFrame, window: FocusTimeWindow) -> None: - payload: Final = self._serializer.serialize(frame) + payload: Final = self.serializer.serialize(frame) if not payload: verbose_logger.debug("Focus export: serializer returned empty payload") return - await self._destination.deliver( + await self.destination.deliver( content=payload, time_window=window, - filename=self._build_filename(window), + filename=self.build_filename(window), ) - def _build_filename(self, window: FocusTimeWindow) -> str: - if not self._serializer.extension: + def build_filename(self, window: FocusTimeWindow) -> str: + if not self.serializer.extension: raise ValueError("Serializer must declare a file extension") # Include time window in filename so Vantage (which deduplicates # by filename) doesn't overwrite previous uploads. start_str: Final = window.start_time.strftime("%Y%m%dT%H%M%SZ") end_str: Final = window.end_time.strftime("%Y%m%dT%H%M%SZ") - return f"usage_{start_str}_{end_str}.{self._serializer.extension}" + return f"usage_{start_str}_{end_str}.{self.serializer.extension}" + + _build_filename = build_filename @staticmethod - def _sum_column(frame: pl.DataFrame, column: str) -> float: + def sum_column(frame: pl.DataFrame, column: str) -> float: if frame.is_empty() or column not in frame.columns: return 0.0 - value: Final = frame.select(pl.col(column).sum().alias("sum")).row(0)[0] + value: Final[object] = frame.select(pl.col(column).sum().alias("sum")).row(0)[0] if value is None: return 0.0 - return float(value) + return float(cast(SupportsFloat | SupportsIndex | str | bytes | bytearray, value)) + + _sum_column: Final[Callable[[pl.DataFrame, str], float]] = sum_column @staticmethod - def _count_unique(frame: pl.DataFrame, column: str) -> int: + def count_unique(frame: pl.DataFrame, column: str) -> int: if frame.is_empty() or column not in frame.columns: return 0 - value: Final = frame.select(pl.col(column).n_unique().alias("unique")).row(0)[0] + value: Final[object] = frame.select(pl.col(column).n_unique().alias("unique")).row(0)[0] if value is None: return 0 - return int(value) + return int(cast(SupportsInt | SupportsIndex | str | bytes | bytearray, value)) + + _count_unique: Final[Callable[[pl.DataFrame, str], int]] = count_unique diff --git a/litellm/integrations/gitlab/gitlab_prompt_manager.py b/litellm/integrations/gitlab/gitlab_prompt_manager.py index 817d280074f..113dea541fb 100644 --- a/litellm/integrations/gitlab/gitlab_prompt_manager.py +++ b/litellm/integrations/gitlab/gitlab_prompt_manager.py @@ -120,17 +120,19 @@ class GitLabTemplateManager: ) if self.prompt_id: - self._load_prompt_from_gitlab(self.prompt_id) + self.load_prompt_from_gitlab(self.prompt_id) # ---------- path helpers ---------- - def _id_to_repo_path(self, prompt_id: str) -> str: + def id_to_repo_path(self, prompt_id: str) -> str: """Map a prompt_id to a repo path (respects prompts_path and adds .prompt).""" prompt_id = decode_prompt_id(prompt_id) if self.prompts_path: return f"{self.prompts_path}/{prompt_id}.prompt" return f"{prompt_id}.prompt" + _id_to_repo_path = id_to_repo_path + def _repo_path_to_id(self, repo_path: str) -> str: """ Map a repo path like 'prompts/chat/greeting.prompt' to an ID relative @@ -144,11 +146,11 @@ class GitLabTemplateManager: # ---------- loading ---------- - def _load_prompt_from_gitlab(self, prompt_id: str, *, ref: str | None = None) -> None: + def load_prompt_from_gitlab(self, prompt_id: str, *, ref: str | None = None) -> None: """Load a specific .prompt file from GitLab (scoped under prompts_path if set).""" try: # prompt_id = decode_prompt_id(prompt_id) - file_path: Final = self._id_to_repo_path(prompt_id) + file_path: Final = self.id_to_repo_path(prompt_id) prompt_content: Final = self.gitlab_client.get_file_content(file_path, ref=ref) if prompt_content: template: Final = self._parse_prompt_file(prompt_content, prompt_id) @@ -156,6 +158,8 @@ class GitLabTemplateManager: except Exception as e: raise Exception(f"Failed to load prompt '{encode_prompt_id(prompt_id)}' from GitLab: {e}") + _load_prompt_from_gitlab = load_prompt_from_gitlab + def load_all_prompts(self, *, recursive: bool = True) -> list[str]: """ Eagerly load all .prompt files from prompts_path. Returns loaded IDs. @@ -164,7 +168,7 @@ class GitLabTemplateManager: loaded: Final[list[str]] = [] for pid in files: if pid not in self.prompts: - self._load_prompt_from_gitlab(pid) + self.load_prompt_from_gitlab(pid) loaded.append(pid) return loaded @@ -333,7 +337,7 @@ class GitLabPromptManager(CustomPromptManagement): ref: str | None = None, ) -> tuple[str, dict[str, Any]]: if prompt_id not in self.prompt_manager.prompts: - self.prompt_manager._load_prompt_from_gitlab(prompt_id, ref=ref) + self.prompt_manager.load_prompt_from_gitlab(prompt_id, ref=ref) template: Final = self.prompt_manager.get_template(prompt_id) if not template: @@ -506,7 +510,7 @@ class GitLabPromptManager(CustomPromptManagement): if hasattr(dynamic_callback_params, "extra") else None ) - self.prompt_manager._load_prompt_from_gitlab(decoded_id, ref=git_ref) + self.prompt_manager.load_prompt_from_gitlab(decoded_id, ref=git_ref) rendered_prompt, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables) @@ -689,17 +693,17 @@ class GitLabPromptCache: for pid in ids: # Ensure template is loaded into TemplateManager if pid not in self.template_manager.prompts: - self.template_manager._load_prompt_from_gitlab(pid) + self.template_manager.load_prompt_from_gitlab(pid) tmpl = self.template_manager.get_template(pid) if tmpl is None: # If something raced/failed, try once more - self.template_manager._load_prompt_from_gitlab(pid) + self.template_manager.load_prompt_from_gitlab(pid) tmpl = self.template_manager.get_template(pid) if tmpl is None: continue - file_path = self.template_manager._id_to_repo_path(pid) # "prompts/chat/..../file.prompt" + file_path = self.template_manager.id_to_repo_path(pid) # "prompts/chat/..../file.prompt" entry = self._template_to_json(pid, tmpl) self._by_file[file_path] = entry @@ -758,7 +762,7 @@ class GitLabPromptCache: return { "id": prompt_id, # e.g. "greet/hi" - "path": self.template_manager._id_to_repo_path(prompt_id), # e.g. "prompts/chat/greet/hi.prompt" + "path": self.template_manager.id_to_repo_path(prompt_id), # e.g. "prompts/chat/greet/hi.prompt" "content": tmpl.content, # rendered content (without frontmatter) "metadata": md, # parsed frontmatter "model": model, diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 055819df86c..126ea1da728 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -1023,7 +1023,7 @@ class LangFuseLogger: _cache_key = _hidden_params.get("cache_key", None) if _cache_key is None and litellm.cache is not None: # fallback to using "preset_cache_key" - _preset_cache_key: Final = litellm.cache._get_preset_cache_key_from_kwargs(**kwargs) # pyright: ignore[reportPrivateUsage] # kwargs-ok: no public preset-cache-key accessor + _preset_cache_key: Final = litellm.cache.get_preset_cache_key_from_kwargs(**kwargs) _cache_key = _preset_cache_key tags.append(f"cache_key:{_cache_key}") return tags diff --git a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py index 4c2f75bb4d7..6e73387326a 100644 --- a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py +++ b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py @@ -116,7 +116,7 @@ class MavvrikFocusLogger(FocusLogger): """Export with Mavvrik row cap applied when no explicit limit is passed.""" effective_limit: Final = limit if limit is not None else self._max_rows engine: Final = self._ensure_engine() - data: Final = await engine._database.get_usage_data( + data: Final = await engine.database.get_usage_data( limit=effective_limit, start_time_utc=window.start_time, end_time_utc=window.end_time, @@ -134,13 +134,13 @@ class MavvrikFocusLogger(FocusLogger): if data.is_empty(): verbose_proxy_logger.debug("Mavvrik FOCUS export: no usage data for window %s", window) else: - normalized: Final = engine._transformer.transform(data) + normalized: Final = engine.transformer.transform(data) if not normalized.is_empty(): - payload = engine._serializer.serialize(normalized) - await engine._destination.deliver( + payload = engine.serializer.serialize(normalized) + await engine.destination.deliver( content=payload or b"", time_window=window, - filename=engine._build_filename(window), + filename=engine.build_filename(window), ) # Maximum number of days to catch up in a single run. Prevents runaway @@ -165,7 +165,7 @@ class MavvrikFocusLogger(FocusLogger): FocusMavvrikDestination, ) - destination: Final = engine._destination + destination: Final = engine.destination if not isinstance(destination, FocusMavvrikDestination): await super()._run_scheduled_export() return diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 6f16b124e6a..90383ddb60d 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -181,7 +181,7 @@ class OTELMetricAttributeFilter: exclude_list: list[str] | None = None -def _build_metric_attribute_filter(value: object) -> OTELMetricAttributeFilter: +def build_metric_attribute_filter(value: object) -> OTELMetricAttributeFilter: if isinstance(value, OTELMetricAttributeFilter): return value if not isinstance(value, dict): @@ -195,7 +195,10 @@ def _build_metric_attribute_filter(value: object) -> OTELMetricAttributeFilter: ) -def _resolve_metric_attribute_filter( +_build_metric_attribute_filter = build_metric_attribute_filter + + +def resolve_metric_attribute_filter( attributes: OTELMetricAttributeFilter | None, ) -> tuple[frozenset[str] | None, frozenset[str] | None]: if attributes is None: @@ -220,6 +223,9 @@ def _resolve_metric_attribute_filter( ) +_resolve_metric_attribute_filter = resolve_metric_attribute_filter + + def _provider_label(custom_llm_provider: object) -> str | None: """The provider label for one call's metrics and events, or None when the call carries no provider. @@ -408,7 +414,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): if metadata_keys_override is not None: config.baggage_metadata_keys = _normalize_team_metadata_keys(metadata_keys_override) if metric_attributes_override is not None: - config.attributes = _build_metric_attribute_filter(metric_attributes_override) + config.attributes = build_metric_attribute_filter(metric_attributes_override) self.config = config self.callback_name = callback_name @@ -1643,11 +1649,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {} raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None if raw is not None: - attributes = _build_metric_attribute_filter(raw) + attributes = build_metric_attribute_filter(raw) ( self._metric_attr_include, self._metric_attr_exclude, - ) = _resolve_metric_attribute_filter(attributes) + ) = resolve_metric_attribute_filter(attributes) self._metric_attr_filter_resolved = True def _filter_metric_attributes(self, attrs: Mapping[str, str | None]) -> dict[str, str]: diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index 76b23467679..7ec3711ee65 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -21,8 +21,8 @@ from litellm._logging import verbose_logger from litellm.integrations.opentelemetry import ( METRIC_METADATA_KEYS, TOKEN_TYPE_ATTRIBUTE, - _build_metric_attribute_filter, - _resolve_metric_attribute_filter, + build_metric_attribute_filter, + resolve_metric_attribute_filter, ) from litellm.integrations.otel.model.metadata import time_to_first_chunk_seconds from litellm.integrations.otel.model.semconv import ( @@ -324,13 +324,13 @@ class GenAIMetricRecorder: otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {} raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None if raw is not None: - attributes = _build_metric_attribute_filter(raw) + attributes = build_metric_attribute_filter(raw) # A bad filter (include_list + exclude_list both set, an unfilterable name) # raises here; the caller (logger._record_metrics) surfaces it once at ERROR # so the operator-fixable config error is visible. Not cached on the raise # path -- _filter_resolved stays False -- so a corrected config takes effect # without reconstructing the recorder. - self._include, self._exclude = _resolve_metric_attribute_filter(attributes) + self._include, self._exclude = resolve_metric_attribute_filter(attributes) self._filter_resolved = True self._warn_about_metric_ineligible_names() diff --git a/litellm/integrations/otel/presets/weave.py b/litellm/integrations/otel/presets/weave.py index 856e784460a..a6eb8bb8cdf 100644 --- a/litellm/integrations/otel/presets/weave.py +++ b/litellm/integrations/otel/presets/weave.py @@ -12,7 +12,7 @@ from litellm.integrations.otel.presets.utils import ( ensure_mappers, ) from litellm.integrations.weave.weave_otel import ( - _get_weave_authorization_header, + get_weave_authorization_header, get_weave_otel_config, ) from litellm.types.utils import StandardCallbackDynamicParams @@ -58,7 +58,7 @@ def weave_dynamic_headers(params: StandardCallbackDynamicParams) -> dict[str, st headers: Final[dict[str, str]] = {} api_key: Final = params.get("wandb_api_key") if api_key: - headers["Authorization"] = _get_weave_authorization_header(api_key=api_key) + headers["Authorization"] = get_weave_authorization_header(api_key=api_key) project_id: Final = params.get("weave_project_id") if project_id: headers["project_id"] = project_id diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 4366f785803..ede75381277 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -32,7 +32,7 @@ from litellm.exceptions import ( from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.prometheus_helpers import ( PrometheusLabelFactoryContext, - _get_cached_end_user_id_for_cost_tracking, + get_cached_end_user_id_for_cost_tracking, ) from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import ( BoundedPrometheusSeriesTracker, @@ -65,8 +65,8 @@ from litellm.repositories.user_repository import UserRepository from litellm.types.guardrails import GuardrailEventHooks from litellm.types.integrations.prometheus import * from litellm.types.integrations.prometheus import ( - _sanitize_prometheus_label_name, - _sanitize_prometheus_label_value, + sanitize_prometheus_label_name, + sanitize_prometheus_label_value, validate_prometheus_deployment_and_latency_caller_identity, ) from litellm.types.proxy.carried_budget_state import ( @@ -1069,10 +1069,10 @@ class PrometheusLogger(CustomLogger): builtin_labels: Final = frozenset(label.value for label in UserAPIKeyLabelNames) custom_metadata_labels: Final = frozenset( - _sanitize_prometheus_label_name(label) for label in litellm.custom_prometheus_metadata_labels + sanitize_prometheus_label_name(label) for label in litellm.custom_prometheus_metadata_labels ) custom_tag_labels: Final = frozenset( - _sanitize_prometheus_label_name(f"tag_{tag}") for tag in litellm.custom_prometheus_tags + sanitize_prometheus_label_name(f"tag_{tag}") for tag in litellm.custom_prometheus_tags ) return builtin_labels | _NON_ENUM_METRIC_LABELS | custom_metadata_labels | custom_tag_labels @@ -1508,7 +1508,7 @@ class PrometheusLogger(CustomLogger): model: Final = kwargs.get("model", "") litellm_params: Final = kwargs.get("litellm_params", {}) or {} _metadata: Final = litellm_params.get("metadata") or {} - get_end_user_id_for_cost_tracking: Final = _get_cached_end_user_id_for_cost_tracking() + get_end_user_id_for_cost_tracking: Final = get_cached_end_user_id_for_cost_tracking() end_user_id: Final = get_end_user_id_for_cost_tracking(litellm_params, service_type="prometheus") user_id: Final = standard_logging_payload["metadata"]["user_api_key_user_id"] @@ -2523,7 +2523,7 @@ class PrometheusLogger(CustomLogger): model: Final = kwargs.get("model", "") litellm_params: Final = kwargs.get("litellm_params", {}) or {} - get_end_user_id_for_cost_tracking: Final = _get_cached_end_user_id_for_cost_tracking() + get_end_user_id_for_cost_tracking: Final = get_cached_end_user_id_for_cost_tracking() end_user_id: Final = get_end_user_id_for_cost_tracking(litellm_params, service_type="prometheus") user_id: Final = standard_logging_payload["metadata"]["user_api_key_user_id"] @@ -2775,7 +2775,7 @@ class PrometheusLogger(CustomLogger): status_code: Final = self._extract_status_code(exception=original_exception) try: - _tags: Final = StandardLoggingPayloadSetup._get_request_tags( + _tags: Final = StandardLoggingPayloadSetup.get_request_tags( litellm_params=request_data, proxy_server_request=request_data.get("proxy_server_request", {}), ) @@ -3725,11 +3725,11 @@ class PrometheusLogger(CustomLogger): increment metric when litellm.Router / load balancing logic places a deployment in cool down """ self.litellm_deployment_cooled_down.labels( - _sanitize_prometheus_label_value(litellm_model_name), - _sanitize_prometheus_label_value(model_id), - _sanitize_prometheus_label_value(api_base), - _sanitize_prometheus_label_value(api_provider), - _sanitize_prometheus_label_value(exception_status), + sanitize_prometheus_label_value(litellm_model_name), + sanitize_prometheus_label_value(model_id), + sanitize_prometheus_label_value(api_base), + sanitize_prometheus_label_value(api_provider), + sanitize_prometheus_label_value(exception_status), ).inc() def increment_callback_logging_failure( @@ -4750,17 +4750,17 @@ def _prometheus_labels_from_context( ctx: PrometheusLabelFactoryContext, ) -> dict[str, str | None]: filtered_labels: Final[dict[str, str | None]] = { - label: ctx._sanitized_enum[label] for label in supported_enum_labels if label in ctx._sanitized_enum + label: ctx.sanitized_enum[label] for label in supported_enum_labels if label in ctx.sanitized_enum } if UserAPIKeyLabelNames.END_USER.value in filtered_labels: filtered_labels[UserAPIKeyLabelNames.END_USER.value] = ctx.get_resolved_end_user() - for sk, val in ctx._custom_by_sanitized_key.items(): + for sk, val in ctx.custom_by_sanitized_key.items(): if sk in supported_enum_labels: filtered_labels[sk] = val - for k, v in ctx._tag_labels.items(): + for k, v in ctx.tag_labels.items(): if k in supported_enum_labels: filtered_labels[k] = v @@ -4797,13 +4797,13 @@ def prometheus_label_factory( # Filter supported labels and sanitize values to prevent breaking # the Prometheus text format (e.g. U+2028 Line Separator in label values) filtered_labels: Final = { - label: _sanitize_prometheus_label_value(value) + label: sanitize_prometheus_label_value(value) for label, value in enum_dict.items() if label in supported_enum_labels } if UserAPIKeyLabelNames.END_USER.value in filtered_labels: - get_end_user_id_for_cost_tracking: Final = _get_cached_end_user_id_for_cost_tracking() + get_end_user_id_for_cost_tracking: Final = get_cached_end_user_id_for_cost_tracking() filtered_labels["end_user"] = get_end_user_id_for_cost_tracking( litellm_params={"user_api_key_end_user_id": enum_values.end_user}, @@ -4813,16 +4813,16 @@ def prometheus_label_factory( if enum_values.custom_metadata_labels is not None: for key, value in enum_values.custom_metadata_labels.items(): # check sanitized key - sanitized_key = _sanitize_prometheus_label_name(key) + sanitized_key = sanitize_prometheus_label_name(key) if sanitized_key in supported_enum_labels: - filtered_labels[sanitized_key] = _sanitize_prometheus_label_value(value) + filtered_labels[sanitized_key] = sanitize_prometheus_label_value(value) # Add custom tags if configured if enum_values.tags is not None: custom_tag_labels: Final = get_custom_labels_from_tags(enum_values.tags) for key, value in custom_tag_labels.items(): if key in supported_enum_labels: - filtered_labels[key] = _sanitize_prometheus_label_value(value) + filtered_labels[key] = sanitize_prometheus_label_value(value) for label in supported_enum_labels: if label not in filtered_labels: @@ -4919,7 +4919,7 @@ def _tag_matches_wildcard_configured_pattern(tags: Sequence[str], configured_tag from litellm.router_utils.pattern_match_deployments import PatternMatchRouter pattern_router: Final = PatternMatchRouter() - regex_pattern: Final = pattern_router._pattern_to_regex(configured_tag) + regex_pattern: Final = pattern_router.pattern_to_regex(configured_tag) return any(re.match(pattern=regex_pattern, string=tag) for tag in tags) @@ -4945,7 +4945,7 @@ def get_custom_labels_from_tags(tags: Sequence[str]) -> dict[str, str]: } """ - from litellm.types.integrations.prometheus import _sanitize_prometheus_label_name + from litellm.types.integrations.prometheus import sanitize_prometheus_label_name configured_tags: Final = litellm.custom_prometheus_tags if configured_tags is None or len(configured_tags) == 0: @@ -4954,7 +4954,7 @@ def get_custom_labels_from_tags(tags: Sequence[str]) -> dict[str, str]: result: Final[dict[str, str]] = {} for configured_tag in configured_tags: - label_name = _sanitize_prometheus_label_name(f"tag_{configured_tag}") + label_name = sanitize_prometheus_label_name(f"tag_{configured_tag}") # Check for exact match first (backwards compatibility) if configured_tag in tags: diff --git a/litellm/integrations/prometheus_helpers/__init__.py b/litellm/integrations/prometheus_helpers/__init__.py index 2e6a350e3c2..5c4cba6a2eb 100644 --- a/litellm/integrations/prometheus_helpers/__init__.py +++ b/litellm/integrations/prometheus_helpers/__init__.py @@ -6,18 +6,26 @@ Helpers for the Prometheus integration (extracted to keep ``prometheus.py`` smal from __future__ import annotations -from typing import Final, cast +from typing import Final, Literal, Protocol, cast from litellm.types.integrations.prometheus import ( UserAPIKeyLabelValues, - _sanitize_prometheus_label_name, - _sanitize_prometheus_label_value, + sanitize_prometheus_label_name, + sanitize_prometheus_label_value, ) _get_end_user_id_for_cost_tracking = None -def _get_cached_end_user_id_for_cost_tracking(): +class _EndUserIdGetter(Protocol): + def __call__( + self, + litellm_params: dict[str, object], + service_type: Literal["litellm_logging", "prometheus"] = "litellm_logging", + ) -> str | None: ... + + +def get_cached_end_user_id_for_cost_tracking() -> _EndUserIdGetter: """ Get cached get_end_user_id_for_cost_tracking function. Lazy imports on first call to avoid loading utils.py at import time (60MB saved). @@ -31,6 +39,9 @@ def _get_cached_end_user_id_for_cost_tracking(): return _get_end_user_id_for_cost_tracking +_get_cached_end_user_id_for_cost_tracking = get_cached_end_user_id_for_cost_tracking + + class PrometheusLabelFactoryContext: """ Precomputes per-request label inputs so prometheus_label_factory can subset @@ -38,11 +49,11 @@ class PrometheusLabelFactoryContext: """ __slots__ = ( - "_custom_by_sanitized_key", "_resolved_end_user", - "_sanitized_enum", - "_tag_labels", + "custom_by_sanitized_key", "enum_values", + "sanitized_enum", + "tag_labels", ) _END_USER_NOT_COMPUTED = object() @@ -50,27 +61,51 @@ class PrometheusLabelFactoryContext: def __init__(self, enum_values: UserAPIKeyLabelValues) -> None: self.enum_values = enum_values enum_dict: Final = enum_values.model_dump() - self._sanitized_enum: dict[str, str | None] = { - k: _sanitize_prometheus_label_value(v) for k, v in enum_dict.items() + self.sanitized_enum: dict[str, str | None] = { + k: sanitize_prometheus_label_value(v) for k, v in enum_dict.items() } - self._custom_by_sanitized_key: dict[str, str | None] = {} + self.custom_by_sanitized_key: dict[str, str | None] = {} if enum_values.custom_metadata_labels is not None: for key, value in enum_values.custom_metadata_labels.items(): - sk = _sanitize_prometheus_label_name(key) - self._custom_by_sanitized_key[sk] = _sanitize_prometheus_label_value(value) - self._tag_labels: dict[str, str | None] = {} + sk = sanitize_prometheus_label_name(key) + self.custom_by_sanitized_key[sk] = sanitize_prometheus_label_value(value) + self.tag_labels: dict[str, str | None] = {} if enum_values.tags is not None: # Late import avoids circular import: ``prometheus`` imports this module. from litellm.integrations.prometheus import get_custom_labels_from_tags for k, v in get_custom_labels_from_tags(enum_values.tags).items(): - self._tag_labels[k] = _sanitize_prometheus_label_value(v) + self.tag_labels[k] = sanitize_prometheus_label_value(v) # Use a dedicated sentinel so `None` can be cached as a computed result. self._resolved_end_user: object = self._END_USER_NOT_COMPUTED + @property + def _custom_by_sanitized_key(self) -> dict[str, str | None]: + return self.custom_by_sanitized_key + + @_custom_by_sanitized_key.setter + def _custom_by_sanitized_key(self, value: dict[str, str | None]) -> None: + self.custom_by_sanitized_key = value + + @property + def _sanitized_enum(self) -> dict[str, str | None]: + return self.sanitized_enum + + @_sanitized_enum.setter + def _sanitized_enum(self, value: dict[str, str | None]) -> None: + self.sanitized_enum = value + + @property + def _tag_labels(self) -> dict[str, str | None]: + return self.tag_labels + + @_tag_labels.setter + def _tag_labels(self, value: dict[str, str | None]) -> None: + self.tag_labels = value + def get_resolved_end_user(self) -> str | None: if self._resolved_end_user is self._END_USER_NOT_COMPUTED: - fn: Final = _get_cached_end_user_id_for_cost_tracking() + fn: Final = get_cached_end_user_id_for_cost_tracking() self._resolved_end_user = fn( litellm_params={"user_api_key_end_user_id": self.enum_values.end_user}, service_type="prometheus", diff --git a/litellm/integrations/prometheus_services.py b/litellm/integrations/prometheus_services.py index 7f76f6db792..d0db4888c50 100644 --- a/litellm/integrations/prometheus_services.py +++ b/litellm/integrations/prometheus_services.py @@ -113,7 +113,9 @@ class PrometheusServicesLogger: """ Helper function to get a metric from the registry by name. """ - return self.REGISTRY._names_to_collectors.get(metric_name) + return self.REGISTRY._names_to_collectors.get( # pyright: ignore[reportPrivateUsage] # Registry lookup has no public API + metric_name + ) def create_histogram(self, service: str, type_of_request: str): metric_name: Final = f"litellm_{service}_{type_of_request}" @@ -196,7 +198,7 @@ class PrometheusServicesLogger: labels=payload.service.value, amount=payload.duration, ) - elif isinstance(obj, self.Counter) and "total_requests" in obj._name: + elif isinstance(obj, self.Counter) and "total_requests" in obj._name: # pyright: ignore[reportPrivateUsage] # Metric names have no public accessor self.increment_counter( counter=obj, labels=payload.service.value, @@ -233,7 +235,7 @@ class PrometheusServicesLogger: labels=payload.service.value, amount=payload.duration, ) - elif isinstance(obj, self.Counter) and "total_requests" in obj._name: + elif isinstance(obj, self.Counter) and "total_requests" in obj._name: # pyright: ignore[reportPrivateUsage] # Metric names have no public accessor self.increment_counter( counter=obj, labels=payload.service.value, @@ -262,7 +264,7 @@ class PrometheusServicesLogger: for obj in prom_objects: # increment both failed and total requests if isinstance(obj, self.Counter): - if "failed_requests" in obj._name: + if "failed_requests" in obj._name: # pyright: ignore[reportPrivateUsage] # Metric names have no public accessor self.increment_counter( counter=obj, labels=payload.service.value, diff --git a/litellm/integrations/weave/weave_otel.py b/litellm/integrations/weave/weave_otel.py index 50289263f38..2bff3682de7 100644 --- a/litellm/integrations/weave/weave_otel.py +++ b/litellm/integrations/weave/weave_otel.py @@ -106,7 +106,7 @@ def _set_weave_specific_attributes(span: Span, kwargs: Mapping[str, Any], respon safe_set_attribute(span, OpenInferenceSpanAttributes.OUTPUT_VALUE, safe_dumps(output_dict)) -def _get_weave_authorization_header(api_key: str) -> str: +def get_weave_authorization_header(api_key: str) -> str: """ Get the authorization header for Weave OpenTelemetry. @@ -117,6 +117,9 @@ def _get_weave_authorization_header(api_key: str) -> str: return f"Basic {auth_header}" +_get_weave_authorization_header = get_weave_authorization_header + + def weave_otel_endpoint(host: str | None) -> str: """The OTLP traces endpoint for a self-managed ``host``, else Weave cloud.""" if not host: @@ -155,7 +158,7 @@ def get_weave_otel_config() -> WeaveOtelConfig: verbose_logger.debug("Using Weave OTEL endpoint: %s", endpoint) # Weave uses Basic auth with format: api: - auth_header: Final = _get_weave_authorization_header(api_key=api_key) + auth_header: Final = get_weave_authorization_header(api_key=api_key) otlp_auth_headers: Final = f"Authorization={auth_header},project_id={project_id}" # Set standard OTEL environment variables @@ -320,7 +323,7 @@ class WeaveOtelLogger(OpenTelemetry): dynamic_weave_project_id: Final = standard_callback_dynamic_params.get("weave_project_id") if dynamic_wandb_api_key: - auth_header: Final = _get_weave_authorization_header( + auth_header: Final = get_weave_authorization_header( api_key=dynamic_wandb_api_key, ) dynamic_headers["Authorization"] = auth_header diff --git a/litellm/interactions/streaming_iterator.py b/litellm/interactions/streaming_iterator.py index 48f2ed68457..64b599edfe2 100644 --- a/litellm/interactions/streaming_iterator.py +++ b/litellm/interactions/streaming_iterator.py @@ -72,7 +72,7 @@ class BaseInteractionsAPIStreamingIterator: return None # Handle SSE format (data: {...}) - stripped_chunk: Final = CustomStreamWrapper._strip_sse_data_from_chunk(chunk) + stripped_chunk: Final = CustomStreamWrapper.strip_sse_data_from_chunk(chunk) if stripped_chunk is None: return None diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 76ea2c25159..de0e096a764 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -5,7 +5,7 @@ import logging import re from collections.abc import Collection, Iterable, Mapping from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast import httpx from pydantic import TypeAdapter, ValidationError @@ -427,30 +427,43 @@ def reconstruct_model_name( # Helper functions used for OTEL logging -def _get_parent_otel_span_from_kwargs( - kwargs: dict | None = None, +def get_parent_otel_span_from_kwargs( + kwargs: dict[str, object] | None = None, ) -> Span | None: try: if kwargs is None: return None litellm_params: Final = kwargs.get("litellm_params") - _metadata: Final = kwargs.get("metadata") or {} - if "litellm_parent_otel_span" in _metadata: - return _metadata["litellm_parent_otel_span"] + metadata: Final = cast( # cast-ok: metadata is caller-provided request data + Mapping[str, object], kwargs.get("metadata") or {} + ) + if "litellm_parent_otel_span" in metadata: + return cast( # cast-ok: tracing metadata crosses an external boundary + Span | None, metadata["litellm_parent_otel_span"] + ) elif ( litellm_params is not None - and litellm_params.get("metadata") is not None - and "litellm_parent_otel_span" in litellm_params.get("metadata", {}) + and cast(Mapping[str, object], litellm_params).get("metadata") is not None + and "litellm_parent_otel_span" + in cast( + Mapping[str, object], + cast(Mapping[str, object], litellm_params).get("metadata", {}), + ) ): - return litellm_params["metadata"]["litellm_parent_otel_span"] + typed_litellm_params: Final = cast(Mapping[str, object], litellm_params) + litellm_metadata: Final = cast(Mapping[str, object], typed_litellm_params["metadata"]) + return cast(Span | None, litellm_metadata["litellm_parent_otel_span"]) elif "litellm_parent_otel_span" in kwargs: - return kwargs["litellm_parent_otel_span"] + return cast(Span | None, kwargs["litellm_parent_otel_span"]) return None except Exception as e: verbose_logger.exception("Error in _get_parent_otel_span_from_kwargs: " + str(e)) return None +_get_parent_otel_span_from_kwargs = get_parent_otel_span_from_kwargs + + def process_response_headers( response_headers: httpx.Headers | dict, preserve_litellm_internal_headers: bool = False, diff --git a/litellm/litellm_core_utils/dd_tracing.py b/litellm/litellm_core_utils/dd_tracing.py index bbbaa843622..a4db3bcd5c2 100644 --- a/litellm/litellm_core_utils/dd_tracing.py +++ b/litellm/litellm_core_utils/dd_tracing.py @@ -57,11 +57,14 @@ def _should_use_dd_tracer(): return get_secret_bool("USE_DDTRACE", False) is True -def _should_use_dd_profiler(): +def should_use_dd_profiler() -> bool: """Returns True if `USE_DDPROFILER` is set to True in .env""" return get_secret_bool("USE_DDPROFILER", False) is True +_should_use_dd_profiler = should_use_dd_profiler + + # Initialize tracer should_use_dd_tracer: Final = _should_use_dd_tracer() tracer: NullTracer | DD_TRACER = NullTracer() diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 0fdfb301291..93c2f869c0c 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -9,13 +9,13 @@ from typing import Final, Protocol, cast import httpx import litellm -from litellm._logging import _ENABLE_SECRET_REDACTION, _redact_string, verbose_logger +from litellm._logging import _ENABLE_SECRET_REDACTION, redact_string, verbose_logger from litellm.litellm_core_utils.bug_report import ( bug_report_notice, build_bug_report, should_report_bug, ) -from litellm.litellm_core_utils.secret_redaction import redact_string +from litellm.litellm_core_utils.secret_redaction import redact_string as redact_secret_string from litellm.types.utils import LlmProviders from ..exceptions import ( @@ -2128,7 +2128,7 @@ def _map_azure_exception( else: # if no status code then it is an APIConnectionError: https://github.com/openai/openai-python#handling-errors raise APIConnectionError( - message=f"{exception_provider} APIConnectionError - {message}\n{_redact_string(traceback.format_exc())}", + message=f"{exception_provider} APIConnectionError - {message}\n{redact_string(traceback.format_exc())}", llm_provider="azure", model=model, litellm_debug_info=extra_information, @@ -2373,18 +2373,20 @@ def exception_type( "\033[1;31mGive Feedback / Get Help: https://github.com/BerriAI/litellm/issues/new\033[0m" ) print( # noqa: T201 - "LiteLLM.Info: If you need to debug this error, use `litellm._turn_on_debug()'." + "LiteLLM.Info: If you need to debug this error, use `litellm.turn_on_debug()'." ) print() # noqa: T201 litellm_response_headers: Final = _get_response_headers(original_exception=original_exception) try: - error_str = redact_string(str(original_exception)) if _ENABLE_SECRET_REDACTION else str(original_exception) + error_str = ( + redact_secret_string(str(original_exception)) if _ENABLE_SECRET_REDACTION else str(original_exception) + ) extra_information = "" if model or custom_llm_provider: if hasattr(original_exception, "message"): error_str = ( - redact_string(str(original_exception.message)) + redact_secret_string(str(original_exception.message)) if _ENABLE_SECRET_REDACTION else str(original_exception.message) ) @@ -2425,7 +2427,7 @@ def exception_type( extra_information += f"\nvertex_location: `{_vertex_location}`\n" # on litellm proxy add key name + team to exceptions - extra_information = _add_key_name_and_team_to_alert(request_info=extra_information, metadata=_metadata) + extra_information = add_key_name_and_team_to_alert(request_info=extra_information, metadata=_metadata) except Exception: # DO NOT LET this Block raising the original exception pass @@ -2687,7 +2689,7 @@ def exception_type( else: raise APIConnectionError( message=( - f"{original_exception}\n{_redact_string(traceback.format_exc())}" + f"{original_exception}\n{redact_string(traceback.format_exc())}" + ( "\n" + bug_report_notice( @@ -2726,7 +2728,7 @@ def exception_type( setattr(e, "litellm_response_headers", litellm_response_headers) raise e # it's already mapped raised_exc: Final = APIConnectionError( - message=f"{original_exception}\n{_redact_string(traceback.format_exc())}", + message=f"{original_exception}\n{redact_string(traceback.format_exc())}", llm_provider="", model="", ) @@ -2766,7 +2768,7 @@ def exception_logging( ) -def _add_key_name_and_team_to_alert(request_info: str, metadata: dict) -> str: +def add_key_name_and_team_to_alert(request_info: str, metadata: Mapping[str, object]) -> str: """ Internal helper function for litellm proxy Add the Key Name + Team Name to the error @@ -2783,3 +2785,6 @@ def _add_key_name_and_team_to_alert(request_info: str, metadata: dict) -> str: return request_info except Exception: return request_info + + +_add_key_name_and_team_to_alert = add_key_name_and_team_to_alert diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index 0fe1952db4b..1a48c00bdc2 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -2,7 +2,7 @@ import reprlib from collections.abc import Mapping, MutableMapping from dataclasses import dataclass, fields from types import MappingProxyType -from typing import Final +from typing import Final, cast from pydantic import TypeAdapter, ValidationError @@ -123,17 +123,25 @@ def with_control_options(litellm_params: Mapping[str, object], control: ControlO return {**litellm_params, CONTROL_OPTIONS_KEY: control} -def _get_base_model_from_litellm_call_metadata( - metadata: dict | None, +def get_base_model_from_litellm_call_metadata( + metadata: Mapping[str, object] | None, ) -> str | None: if metadata is None: return None model_info: Final = metadata.get("model_info") if model_info: - return model_info.get("base_model") + model_info_mapping: Final = cast( # cast-ok: model metadata is caller-provided and preserves its mapping shape + Mapping[str, object], model_info + ) + return cast( # cast-ok: model metadata values are caller-provided + str | None, model_info_mapping.get("base_model") + ) return None +_get_base_model_from_litellm_call_metadata = get_base_model_from_litellm_call_metadata + + def get_litellm_params( api_key: str | None = None, force_timeout=600, @@ -229,7 +237,7 @@ def get_litellm_params( "azure_ad_token_provider": azure_ad_token_provider, "user_continue_message": user_continue_message, "base_model": base_model - or (_get_base_model_from_litellm_call_metadata(metadata=metadata) if metadata else None), + or (get_base_model_from_litellm_call_metadata(metadata=metadata) if metadata else None), "litellm_trace_id": litellm_trace_id, "litellm_session_id": litellm_session_id, "hf_model_name": hf_model_name, diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index d4642ae2aad..5a498ead6bc 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -53,7 +53,7 @@ def _endpoint_matches_api_base(endpoint: str, api_base: str) -> bool: return url_path == endpoint_path or url_path.startswith(endpoint_path + "/") -def _is_non_openai_azure_model(model: str) -> bool: +def is_non_openai_azure_model(model: str) -> bool: try: model_name: Final = model.split("/", 1)[1] if model_name in litellm.cohere_chat_models or f"mistral/{model_name}" in litellm.mistral_chat_models: @@ -63,6 +63,9 @@ def _is_non_openai_azure_model(model: str) -> bool: return False +_is_non_openai_azure_model = is_non_openai_azure_model + + def _is_azure_claude_model(model: str) -> bool: """ Check if a model name contains 'claude' (case-insensitive). @@ -195,7 +198,7 @@ def get_llm_provider( # AZURE AI-Studio Logic - Azure AI Studio supports AZURE/Cohere # If User passes azure/command-r-plus -> we should send it to cohere_chat/command-r-plus if model.split("/", 1)[0] == "azure": - if _is_non_openai_azure_model(model): + if is_non_openai_azure_model(model): custom_llm_provider = "openai" return model, custom_llm_provider, dynamic_api_key, api_base diff --git a/litellm/litellm_core_utils/get_model_cost_map.py b/litellm/litellm_core_utils/get_model_cost_map.py index c8f78e9a814..3f7c8984def 100644 --- a/litellm/litellm_core_utils/get_model_cost_map.py +++ b/litellm/litellm_core_utils/get_model_cost_map.py @@ -105,13 +105,15 @@ class GetModelCostMap: cls._loaded_catalog = MappingProxyType({key: MappingProxyType(entry) for key, entry in raw.items()}) @classmethod - def _get_backup_model_count(cls) -> int: + def get_backup_model_count(cls) -> int: """Return the number of models in the local backup (cached int).""" if cls._backup_model_count < 0: backup: Final = cls.load_local_model_cost_map() cls._backup_model_count = _count_model_entries(backup) return cls._backup_model_count + _get_backup_model_count = get_backup_model_count + @staticmethod def _check_is_valid_dict(fetched_map: dict) -> bool: """Check 1: fetched map is a non-empty dict.""" @@ -402,7 +404,7 @@ async def refetch_model_cost_map( return result if not GetModelCostMap.validate_model_cost_map( fetched_map=result.model_cost_map, - backup_model_count=GetModelCostMap._get_backup_model_count(), + backup_model_count=GetModelCostMap.get_backup_model_count(), ): return ModelCostMapReloadUnavailable(reason=f"model cost map from {url} failed integrity validation") _cost_map_source_info.loaded_at = datetime.now(timezone.utc) @@ -600,7 +602,7 @@ def _retry_remote_fetch_in_background( _litellm_import_complete.wait() if not GetModelCostMap.validate_model_cost_map( fetched_map=result.model_cost_map, - backup_model_count=GetModelCostMap._get_backup_model_count(), # pyright: ignore[reportPrivateUsage] # integrity cache + backup_model_count=GetModelCostMap.get_backup_model_count(), ): verbose_logger.warning( "LiteLLM: Fetched model cost map failed integrity check. Using local backup instead. url=%s", @@ -683,7 +685,7 @@ def get_model_cost_map( # Validate using cached count (cheap int comparison, no file I/O) if not GetModelCostMap.validate_model_cost_map( fetched_map=content, - backup_model_count=GetModelCostMap._get_backup_model_count(), + backup_model_count=GetModelCostMap.get_backup_model_count(), ): verbose_logger.warning( "LiteLLM: Fetched model cost map failed integrity check. Using local backup instead. url=%s", diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index fba59b0b983..60e4b7495ed 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -96,7 +96,7 @@ class HealthCheckHelpers: return {} @staticmethod - def _update_model_params_with_health_check_tracking_information( + def update_model_params_with_health_check_tracking_information( model_params: dict, ) -> dict: """ @@ -120,6 +120,10 @@ class HealthCheckHelpers: ) return model_params + _update_model_params_with_health_check_tracking_information = ( + update_model_params_with_health_check_tracking_information + ) + @staticmethod def _get_metadata_for_health_check_call(): """ @@ -219,7 +223,7 @@ class HealthCheckHelpers: get_audio_file_for_health_check, ) from litellm.litellm_core_utils.health_check_utils import DECISIONS_CALL_PARAMS, _filter_model_params - from litellm.realtime_api.main import _realtime_health_check + from litellm.realtime_api.main import realtime_health_check return { "chat": lambda: litellm.acompletion( @@ -264,7 +268,7 @@ class HealthCheckHelpers: query=prompt or "", documents=["my sample text"], ), - "realtime": lambda: _realtime_health_check( + "realtime": lambda: realtime_health_check( model=model, custom_llm_provider=custom_llm_provider, api_base=model_params.get("api_base", None), diff --git a/litellm/litellm_core_utils/health_check_utils.py b/litellm/litellm_core_utils/health_check_utils.py index 5205e9f1284..9b344d08d4b 100644 --- a/litellm/litellm_core_utils/health_check_utils.py +++ b/litellm/litellm_core_utils/health_check_utils.py @@ -2,6 +2,7 @@ Utils used for litellm.ahealth_check() """ +from collections.abc import Mapping from typing import Final from pydantic import TypeAdapter @@ -17,8 +18,8 @@ def _filter_model_params(model_params: dict) -> dict: return {k: v for k, v in model_params.items() if k != "messages"} -def _create_health_check_response(response_headers: dict) -> dict: - response: Final = {} +def create_health_check_response(response_headers: Mapping[str, object]) -> dict[str, object]: + response: Final[dict[str, object]] = {} if response_headers.get("x-ratelimit-remaining-requests", None) is not None: # not provided for dall-e requests response["x-ratelimit-remaining-requests"] = response_headers["x-ratelimit-remaining-requests"] @@ -29,3 +30,6 @@ def _create_health_check_response(response_headers: dict) -> dict: if response_headers.get("x-ms-region", None) is not None: response["x-ms-region"] = response_headers["x-ms-region"] return response + + +_create_health_check_response = create_health_check_response diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 1dd012383cb..9f19872360c 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -11,11 +11,11 @@ import subprocess import sys import time import traceback -from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence +from collections.abc import Awaitable, Callable, Iterable, Iterator, Mapping, Sequence from datetime import datetime as dt_object from functools import lru_cache from types import MappingProxyType, TracebackType -from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Union, cast +from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast from httpx import Response from pydantic import BaseModel, JsonValue @@ -24,8 +24,8 @@ import litellm from litellm import _custom_logger_compatible_callbacks_literal from litellm._internal_context import post_response_phase from litellm._logging import ( - _is_debugging_on, - _redact_string, + is_debugging_on, + redact_string, session_id_var, set_session_id, set_trace_id, @@ -33,7 +33,7 @@ from litellm._logging import ( verbose_logger, ) from litellm._uuid import uuid -from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final +from litellm.batches.batch_utils import batch_cost_is_final, handle_completed_batch from litellm.caching.caching import DualCache from litellm.caching.caching_handler import LLMCachingHandler from litellm.caching.redis_batch import flush_post_call_redis_batches @@ -47,9 +47,9 @@ from litellm.constants import ( from litellm.cost_calculator import ( RealtimeAPITokenUsageProcessor, ResponsesWebSocketTokenUsageProcessor, - _select_model_name_for_cost_calc, get_usage_object, pricing_entry_for_cost_calc, + select_model_name_for_cost_calc, ) from litellm.exceptions import ( BudgetExceededError, @@ -178,7 +178,7 @@ from litellm.types.utils import ( Usage, ) from litellm.types.videos.main import VideoObject -from litellm.utils import _get_base_model_from_metadata, print_verbose +from litellm.utils import get_base_model_from_metadata, print_verbose from ..integrations.argilla import ArgillaLogger from ..integrations.arize.arize_phoenix import ArizePhoenixLogger @@ -575,6 +575,38 @@ class Logging(LiteLLMLoggingBaseClass): baseline_cache_context: "BaselineCacheContext | None" = None baseline_observation: "CapturedBaselineObservation | None" = None + @property + def _defer_async_logging(self) -> bool: + return self.defer_async_logging + + @_defer_async_logging.setter + def _defer_async_logging(self, value: bool) -> None: + self.defer_async_logging = value + + @property + def _enqueue_deferred_logging(self) -> Callable[[], None] | None: + return self.enqueue_deferred_logging + + @_enqueue_deferred_logging.setter + def _enqueue_deferred_logging(self, value: Callable[[], None] | None) -> None: + self.enqueue_deferred_logging = value + + @property + def _llm_caching_handler(self) -> LLMCachingHandler | None: + return self.llm_caching_handler + + @_llm_caching_handler.setter + def _llm_caching_handler(self, value: LLMCachingHandler | None) -> None: + self.llm_caching_handler = value + + @property + def _on_detached_stream_failure(self) -> Callable[[Exception], Awaitable[None]] | None: + return self.on_detached_stream_failure + + @_on_detached_stream_failure.setter + def _on_detached_stream_failure(self, value: Callable[[Exception], Awaitable[None]] | None) -> None: + self.on_detached_stream_failure = value + def __init__( self, model: str, @@ -686,7 +718,7 @@ class Logging(LiteLLMLoggingBaseClass): # once that response is priced, the same way a non-streamed response is logged self.client_facing_stream_model: str | None = None self.zero_cost_warned: bool = False - self._llm_caching_handler: LLMCachingHandler | None = None + self.llm_caching_handler: LLMCachingHandler | None = None # INITIAL LITELLM_PARAMS litellm_params = {} @@ -721,10 +753,10 @@ class Logging(LiteLLMLoggingBaseClass): # Set by proxy request handlers to defer spend-log fire until after # post_call guardrails have run; the @client decorator then stores the # enqueue closure here instead of firing it immediately. - self._defer_async_logging: bool = False - self._enqueue_deferred_logging: Callable[[], None] | None = None + self.defer_async_logging: bool = False + self.enqueue_deferred_logging: Callable[[], None] | None = None self._async_success_scheduled: bool = False - self._on_detached_stream_failure: Callable[[Exception], Awaitable[None]] | None = None + self.on_detached_stream_failure: Callable[[Exception], Awaitable[None]] | None = None self.shadow_eval_request_snapshot: GuardrailRequestSnapshot | None = None def set_response_timing_metrics(self, timing_metrics: Mapping[str, float]) -> None: @@ -826,7 +858,7 @@ class Logging(LiteLLMLoggingBaseClass): trusted-vars channel is the only way credentials reach a per-team logger. """ _trusted_var_prefix: Final = "dd_" if callback == "datadog" else "newrelic_" if callback == "newrelic" else None - _custom_logger_init_args: Final[dict | None] = ( + _custom_logger_init_args: Final[dict[str, object] | None] = ( {k: v for k, v in self._trusted_callback_vars if k.startswith(_trusted_var_prefix)} if _trusted_var_prefix is not None else None @@ -869,8 +901,8 @@ class Logging(LiteLLMLoggingBaseClass): checks if web_search_options in kwargs or tools and sets the corresponding attribute in StandardBuiltInToolsParams """ return StandardBuiltInToolsParams( - web_search_options=StandardBuiltInToolCostTracking._get_web_search_options(kwargs or {}), - file_search=StandardBuiltInToolCostTracking._get_file_search_tool_call(kwargs or {}), + web_search_options=StandardBuiltInToolCostTracking.get_web_search_options(kwargs or {}), + file_search=StandardBuiltInToolCostTracking.get_file_search_tool_call(kwargs or {}), ) def get_router_model_id(self) -> str | None: @@ -930,7 +962,7 @@ class Logging(LiteLLMLoggingBaseClass): } self.litellm_request_debug = litellm_params.get("litellm_request_debug", False) self.logger_fn = litellm_params.get("logger_fn", None) - if _is_debugging_on() or self.litellm_request_debug: + if is_debugging_on() or self.litellm_request_debug: verbose_logger.debug("self.optional_params: %s", self.optional_params) self.model_call_details.update( @@ -1393,12 +1425,12 @@ class Logging(LiteLLMLoggingBaseClass): additional_args=additional_args, data=additional_args.get("complete_input_dict", {}), ) - _metadata["raw_request"] = _redact_string(str(curl_command)) + _metadata["raw_request"] = redact_string(str(curl_command)) except Exception as e: self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict( error=str(e), ) - _metadata["raw_request"] = _redact_string( + _metadata["raw_request"] = redact_string( f"Unable to Log \ raw request: {e}" ) @@ -1495,7 +1527,7 @@ class Logging(LiteLLMLoggingBaseClass): Prints the RAW curl command sent from LiteLLM """ - if _is_debugging_on() or self.litellm_request_debug: + if is_debugging_on() or self.litellm_request_debug: if litellm.json_logs: masked_headers: Final = self._get_masked_headers(headers or {}) masked_api_base: Final = self._get_masked_api_base(str(api_base or "")) @@ -1555,7 +1587,7 @@ class Logging(LiteLLMLoggingBaseClass): Masks the headers of the request sent from LiteLLM """ - return _get_masked_values(headers, ignore_sensitive_values=ignore_sensitive_headers) + return get_masked_values(headers, ignore_sensitive_values=ignore_sensitive_headers) def post_call(self, original_response, input=None, api_key=None, additional_args={}): # Log the exact result from the LLM API, for streaming - log the type of response received @@ -1800,29 +1832,9 @@ class Logging(LiteLLMLoggingBaseClass): if margin_total_amount is not None: self.cost_breakdown["margin_total_amount"] = margin_total_amount - def _response_cost_calculator( + def response_cost_calculator( self, - result: Union[ - ModelResponse, - ModelResponseStream, - EmbeddingResponse, - ImageResponse, - TranscriptionResponse, - TextCompletionResponse, - HttpxBinaryResponseContent, - RerankResponse, - Batch, - FineTuningJob, - ResponsesAPIResponse, - ResponseCompletedEvent, - OpenAIFileObject, - LiteLLMRealtimeStreamLoggingObject, - OpenAIModerationResponse, - "SearchResponse", - DecisionsResponse, - dict, - list, - ], + result: object, cache_hit: bool | None = None, litellm_model_name: str | None = None, router_model_id: str | None = None, @@ -1845,13 +1857,12 @@ class Logging(LiteLLMLoggingBaseClass): return 0.0 transformed_result: Final = self._generate_content_result_as_model_response(result) - if transformed_result is not None: - result = transformed_result + response_result: Final[object] = transformed_result if transformed_result is not None else result priced_result: Final = ( - result.response - if isinstance(result, (ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent)) - else result + response_result.response + if isinstance(response_result, (ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent)) + else response_result ) result_hidden_params: Final = getattr(priced_result, "_hidden_params", None) or MappingProxyType({}) @@ -1881,18 +1892,28 @@ class Logging(LiteLLMLoggingBaseClass): ## RESPONSE COST ## custom_pricing: Final = self._custom_pricing_for(priced_result) - prompt = self._prompt_for_cost_calculation() + prompt: Final = self._prompt_for_cost_calculation() - if cache_hit is None: - cache_hit = self.model_call_details.get("cache_hit", False) + model_value: Final = litellm_model_name or self.model + cost_cache_hit_value: Final = ( + self.model_call_details.get("cache_hit", False) if cache_hit is None else cache_hit + ) + cost_cache_hit: Final[bool | None] = cast( # cast-ok: callback metadata is caller-provided + bool | None, cost_cache_hit_value + ) + provider_value: Final = self.model_call_details.get("custom_llm_provider", None) + cost_custom_llm_provider: Final[str | None] = cast( # cast-ok: callback metadata is caller-provided + str | None, provider_value + ) + base_model: Final = get_base_model_from_metadata(model_call_details=self.model_call_details) try: response_cost_calculator_kwargs: Final = { "response_object": priced_result, - "model": litellm_model_name or self.model, - "cache_hit": cache_hit, - "custom_llm_provider": self.model_call_details.get("custom_llm_provider", None), - "base_model": _get_base_model_from_metadata(model_call_details=self.model_call_details), + "model": model_value, + "cache_hit": cost_cache_hit, + "custom_llm_provider": cost_custom_llm_provider, + "base_model": base_model, "call_type": self.call_type, "optional_params": self.optional_params, "custom_pricing": custom_pricing, @@ -1906,13 +1927,13 @@ class Logging(LiteLLMLoggingBaseClass): else None ), "vertex_location": _resolve_vertex_location_for_cost( - custom_llm_provider=self.model_call_details.get("custom_llm_provider", None), + custom_llm_provider=cost_custom_llm_provider, litellm_params=(self.litellm_params if hasattr(self, "litellm_params") else None), optional_params=self.optional_params, - model=litellm_model_name or self.model, + model=model_value, ), "region_name": _resolve_mantle_region_for_cost( - custom_llm_provider=self.model_call_details.get("custom_llm_provider", None), + custom_llm_provider=cost_custom_llm_provider, litellm_params=self.model_call_details.get("litellm_params"), ), } @@ -1946,12 +1967,12 @@ class Logging(LiteLLMLoggingBaseClass): debug_info = StandardLoggingModelCostFailureDebugInformation( error_str=str(e), traceback_str=_get_traceback_str_for_error(str(e)), - model=response_cost_calculator_kwargs["model"], - cache_hit=response_cost_calculator_kwargs["cache_hit"], - custom_llm_provider=response_cost_calculator_kwargs["custom_llm_provider"], - base_model=response_cost_calculator_kwargs["base_model"], - call_type=response_cost_calculator_kwargs["call_type"], - custom_pricing=response_cost_calculator_kwargs["custom_pricing"], + model=model_value, + cache_hit=cost_cache_hit, + custom_llm_provider=cost_custom_llm_provider, + base_model=base_model, + call_type=self.call_type, + custom_pricing=custom_pricing, ) verbose_logger.debug("response_cost_failure_debug_information: %s", debug_info) self.model_call_details["response_cost_failure_debug_information"] = debug_info @@ -1965,6 +1986,8 @@ class Logging(LiteLLMLoggingBaseClass): return None + _response_cost_calculator = response_cost_calculator + def _record_zero_cost_diagnostic( self, result: object, @@ -2018,7 +2041,7 @@ class Logging(LiteLLMLoggingBaseClass): completion_response=result, custom_llm_provider=custom_llm_provider, custom_pricing=self._custom_pricing_for(result), - base_model=_get_base_model_from_metadata(model_call_details=self.model_call_details), + base_model=get_base_model_from_metadata(model_call_details=self.model_call_details), router_model_id=router_model_id or self.get_router_model_id(), region_name=_resolve_mantle_region_for_cost( custom_llm_provider=custom_llm_provider, @@ -2117,10 +2140,10 @@ class Logging(LiteLLMLoggingBaseClass): | FineTuningJob, cache_hit: bool | None = None, ) -> float | None: - return self._response_cost_calculator(result=result, cache_hit=cache_hit) + return self.response_cost_calculator(result=result, cache_hit=cache_hit) @staticmethod - def _is_sync_litellm_request(litellm_params: dict) -> bool: + def is_sync_litellm_request(litellm_params: Mapping[str, object]) -> bool: """True for sync SDK entrypoints (``completion``), false for async (``acompletion``, etc.).""" return ( litellm_params.get(CallTypes.acompletion.value, False) is not True @@ -2135,6 +2158,8 @@ class Logging(LiteLLMLoggingBaseClass): and litellm_params.get(CallTypes.arealtime.value, False) is not True ) + _is_sync_litellm_request = is_sync_litellm_request + def _is_assembled_stream_success(self, result=None) -> bool: """Final assembled stream export (not a per-chunk success call). @@ -2176,7 +2201,7 @@ class Logging(LiteLLMLoggingBaseClass): self.model_call_details["has_dispatched_final_stream_success"] = True litellm_params: Final = self.model_call_details.get("litellm_params", {}) or {} - sync_sdk: Final = self._is_sync_litellm_request(litellm_params) + sync_sdk: Final = self.is_sync_litellm_request(litellm_params) passthrough: Final = self.call_type == CallTypes.pass_through.value if sync_sdk and not prefer_async_handlers and not passthrough: self.success_handler( @@ -2217,7 +2242,7 @@ class Logging(LiteLLMLoggingBaseClass): """Bill a fully streamed response on the failure log when a post-call hook rejects it.""" usage: Final = getattr(assembled, "usage", None) if isinstance(usage, Usage): - self.record_partial_usage_for_failure(usage, self._response_cost_calculator(result=assembled) or 0.0) + self.record_partial_usage_for_failure(usage, self.response_cost_calculator(result=assembled) or 0.0) async def dispatch_failure_handlers( self, @@ -2237,7 +2262,7 @@ class Logging(LiteLLMLoggingBaseClass): the request failed). """ litellm_params: Final = self.model_call_details.get("litellm_params", {}) or {} - sync_sdk: Final = self._is_sync_litellm_request(litellm_params) + sync_sdk: Final = self.is_sync_litellm_request(litellm_params) passthrough: Final = self.call_type == CallTypes.pass_through.value if sync_sdk and not prefer_async_handlers and not passthrough: self.failure_handler(exception, traceback_exception) @@ -2314,10 +2339,12 @@ class Logging(LiteLLMLoggingBaseClass): return True - def _update_completion_start_time(self, completion_start_time: datetime.datetime): + def update_completion_start_time(self, completion_start_time: datetime.datetime) -> None: self.completion_start_time = completion_start_time self.model_call_details["completion_start_time"] = self.completion_start_time + _update_completion_start_time = update_completion_start_time + def normalize_logging_result(self, result: object) -> object: """ Some endpoints return a different type of result than what is expected by the logging system. @@ -2434,7 +2461,7 @@ class Logging(LiteLLMLoggingBaseClass): # Do not preserve 0 from failure_handler on intermediate router retries. pass else: - self.model_call_details["response_cost"] = self._response_cost_calculator(result=logging_result) + self.model_call_details["response_cost"] = self.response_cost_calculator(result=logging_result) if not build_logging_payload: return @@ -2482,7 +2509,7 @@ class Logging(LiteLLMLoggingBaseClass): def _transform_usage_objects(self, result): if isinstance(result, ResponsesAPIResponse): result = result.model_copy() - transformed_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(result.usage) + transformed_usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(result.usage) setattr(result, "usage", transformed_usage) if (standard_logging_payload := self.model_call_details.get("standard_logging_object")) is not None: response_dict: Final = result.model_dump() if hasattr(result, "model_dump") else dict(result) @@ -2808,7 +2835,7 @@ class Logging(LiteLLMLoggingBaseClass): standard_logging_object=kwargs.get("standard_logging_object", None), ) litellm_params = self.model_call_details.get("litellm_params", {}) - is_sync_request: Final = self._is_sync_litellm_request(litellm_params) + is_sync_request: Final = self.is_sync_litellm_request(litellm_params) try: ## BUILD COMPLETE STREAMED RESPONSE complete_streaming_response: ( @@ -2827,7 +2854,7 @@ class Logging(LiteLLMLoggingBaseClass): verbose_logger.debug("Logging Details LiteLLM-Success Call streaming complete") self.model_call_details["complete_streaming_response"] = complete_streaming_response self._surface_response_headers_from_result(complete_streaming_response) - self.model_call_details["response_cost"] = self._response_cost_calculator( + self.model_call_details["response_cost"] = self.response_cost_calculator( result=complete_streaming_response ) self._merge_hidden_params_from_response_into_metadata(complete_streaming_response) @@ -3285,7 +3312,7 @@ class Logging(LiteLLMLoggingBaseClass): ) elif should_compute_batch_data: - batch_result: Final = await _handle_completed_batch( + batch_result: Final = await handle_completed_batch( batch=result, custom_llm_provider=self.custom_llm_provider, model_name=self.get_deployment_model_for_cost(), @@ -3351,9 +3378,9 @@ class Logging(LiteLLMLoggingBaseClass): self.model_call_details["response_cost"] = 0.0 else: # check if base_model set on azure - _get_base_model_from_metadata(model_call_details=self.model_call_details) + get_base_model_from_metadata(model_call_details=self.model_call_details) # base_model defaults to None if not set on model_info - self.model_call_details["response_cost"] = self._response_cost_calculator( + self.model_call_details["response_cost"] = self.response_cost_calculator( result=complete_streaming_response ) @@ -3576,7 +3603,7 @@ class Logging(LiteLLMLoggingBaseClass): if self.stream: if "async_complete_streaming_response" in self.model_call_details: print_verbose("DynamoDB Logger: Got Stream Event - Completed Stream Response") - await dynamoLogger._async_log_event( + await dynamoLogger.async_log_event( kwargs=self.model_call_details, response_obj=self.model_call_details["async_complete_streaming_response"], start_time=start_time, @@ -3586,7 +3613,7 @@ class Logging(LiteLLMLoggingBaseClass): else: print_verbose("DynamoDB Logger: Got Stream Event - No complete stream response as yet") else: - await dynamoLogger._async_log_event( + await dynamoLogger.async_log_event( kwargs=self.model_call_details, response_obj=result, start_time=start_time, @@ -3613,7 +3640,7 @@ class Logging(LiteLLMLoggingBaseClass): try: callback_name: Final = self._get_callback_name(callback) - all_callbacks: Final = 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"): @@ -3642,7 +3669,7 @@ class Logging(LiteLLMLoggingBaseClass): self.model_call_details["log_event_type"] = "failed_api_call" self.model_call_details["exception"] = exception self.model_call_details["traceback_exception"] = ( - _redact_string(traceback_exception) if isinstance(traceback_exception, str) else traceback_exception + redact_string(traceback_exception) if isinstance(traceback_exception, str) else traceback_exception ) self.model_call_details["end_time"] = end_time self.model_call_details.setdefault("original_response", None) @@ -3667,7 +3694,7 @@ class Logging(LiteLLMLoggingBaseClass): end_time=end_time, logging_obj=self, status="failure", - error_str=_redact_string(str(exception)), + error_str=redact_string(str(exception)), original_exception=exception, standard_built_in_tools_params=self.standard_built_in_tools_params, ) @@ -3735,7 +3762,7 @@ class Logging(LiteLLMLoggingBaseClass): if not self.should_run_logging(event_type="sync_failure"): # prevent double logging return litellm_params: Final = self.model_call_details.get("litellm_params", {}) - is_sync_request: Final = self._is_sync_litellm_request(litellm_params) + is_sync_request: Final = self.is_sync_litellm_request(litellm_params) try: start_time, end_time = self._failure_handler_helper_fn( @@ -3987,7 +4014,7 @@ class Logging(LiteLLMLoggingBaseClass): self._handle_callback_failure(callback=callback) await flush_post_call_redis_batches() - def _get_trace_id(self, service_name: Literal["langfuse"]) -> str | None: + def get_trace_id(self, service_name: Literal["langfuse"]) -> str | None: """ For the given service (e.g. langfuse), return the trace_id actually logged. @@ -4005,6 +4032,8 @@ class Logging(LiteLLMLoggingBaseClass): return trace_id + _get_trace_id = get_trace_id + def handle_sync_success_callbacks_for_async_calls( self, result: object, @@ -4149,7 +4178,7 @@ class Logging(LiteLLMLoggingBaseClass): ## return unified Usage object if isinstance(result.response.usage, ResponseAPIUsage): set_response_cost_in_hidden_params(result.response, result.response.usage.cost) - transformed_usage: Final = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( + transformed_usage: Final = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage( result.response.usage ) # Set as dict instead of Usage object so model_dump() serializes it correctly @@ -4314,11 +4343,11 @@ class Logging(LiteLLMLoggingBaseClass): model_response: Final = litellm.ModelResponse(id=served_id) model_response.model = self.model usage: Final = getattr(result, "usage", None) - if usage is not None and ResponseAPILoggingUtils._is_response_api_usage(usage): + if usage is not None and ResponseAPILoggingUtils.is_response_api_usage(usage): setattr( model_response, "usage", - ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage), + ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage), ) return model_response @@ -4368,15 +4397,15 @@ class Logging(LiteLLMLoggingBaseClass): return result_copy -def _get_masked_values( - sensitive_object: dict, +def get_masked_values( + sensitive_object: Mapping[str, object], ignore_sensitive_values: bool = False, mask_all_values: bool = False, unmasked_length: int = 4, number_of_asterisks: int | None = 4, _depth: int = 0, _max_depth: int = 20, -) -> dict: +) -> dict[str, object]: """ Internal debugging helper function @@ -4386,7 +4415,7 @@ def _get_masked_values( masked_length: Optional length for the masked portion (number of *). If set, will use exactly this many * regardless of original string length. The total length will be unmasked_length + masked_length. """ - sensitive_keywords: Final = [ + sensitive_keywords: Final = ( "authorization", "token", "key", @@ -4395,14 +4424,17 @@ def _get_masked_values( "credentials", "password", "passwd", - ] + ) def _mask_value(v: object) -> object: if isinstance(v, dict): + typed_value: Final = cast( # cast-ok: request values can contain arbitrary header keys + dict[object, object], v + ) if _depth >= _max_depth: return v - return _get_masked_values( - v, + return get_masked_values( + cast(dict[str, object], typed_value), # cast-ok: preserve dynamic key behavior ignore_sensitive_values=ignore_sensitive_values, mask_all_values=mask_all_values, unmasked_length=unmasked_length, @@ -4429,7 +4461,13 @@ def _get_masked_values( } -def set_callbacks(callback_list, function_id=None): +_get_masked_values = get_masked_values + + +def set_callbacks( + callback_list: Iterable[str | Callable[..., object] | CustomLogger], + function_id: str | None = None, +) -> None: """ Globally sets the callback client """ @@ -4527,14 +4565,14 @@ def set_callbacks(callback_list, function_id=None): def _init_custom_logger_compatible_class( logging_integration: _custom_logger_compatible_callbacks_literal, internal_usage_cache: DualCache | None, - llm_router: object, # expect litellm.Router, but typing errors due to circular import - custom_logger_init_args: dict | None = {}, + llm_router: object, + custom_logger_init_args: dict[str, object] | None = {}, ) -> CustomLogger | None: """ Initialize a custom logger compatible class """ try: - custom_logger_init_args = custom_logger_init_args or {} + custom_logger_init_args_value: Final[dict[str, object]] = custom_logger_init_args or {} if logging_integration == "agentops": # Add AgentOps initialization _v2 = _maybe_construct_otel_v2("agentops", _in_memory_loggers) if _v2 is not None: @@ -4624,10 +4662,10 @@ def _init_custom_logger_compatible_class( return _prometheus_logger elif logging_integration == "datadog": # Check if team-scoped credentials are provided - _dd_api_key: Final = custom_logger_init_args.get("dd_api_key") - _dd_site: Final = custom_logger_init_args.get("dd_site") - _dd_agent_host: Final = custom_logger_init_args.get("dd_agent_host") - _dd_agent_port: Final = custom_logger_init_args.get("dd_agent_port") + _dd_api_key: Final = custom_logger_init_args_value.get("dd_api_key") + _dd_site: Final = custom_logger_init_args_value.get("dd_site") + _dd_agent_host: Final = custom_logger_init_args_value.get("dd_agent_host") + _dd_agent_port: Final = custom_logger_init_args_value.get("dd_agent_port") if _dd_api_key or _dd_site or _dd_agent_host: # Team-scoped credentials: use DynamicLoggingCache for per-credential isolation @@ -4636,7 +4674,9 @@ def _init_custom_logger_compatible_class( ) return DataDogHandler.get_datadog_logger_for_request( - standard_callback_dynamic_params=custom_logger_init_args, + standard_callback_dynamic_params=cast( # cast-ok: callback options come from dynamic config + StandardCallbackDynamicParams, custom_logger_init_args_value + ), in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache, ) @@ -5078,7 +5118,7 @@ def _init_custom_logger_compatible_class( for callback in _in_memory_loggers: if isinstance(callback, pagerduty_loggers.pagerduty): return callback - pagerduty_logger: Final = pagerduty_loggers.pagerduty(**custom_logger_init_args) + pagerduty_logger: Final = pagerduty_loggers.pagerduty(**custom_logger_init_args_value) _in_memory_loggers.append(pagerduty_logger) return pagerduty_logger elif logging_integration == "anthropic_cache_control_hook": @@ -5188,7 +5228,7 @@ def _init_custom_logger_compatible_class( _in_memory_loggers.append(gitlab_logger) return gitlab_logger elif logging_integration == "newrelic": - if custom_logger_init_args.get("newrelic_api_key"): + if custom_logger_init_args_value.get("newrelic_api_key"): # Team-scoped credentials: per-team METRICS logger, isolated per # credential set via DynamicLoggingCache. The trace logger for # this name stays on the global path below. @@ -5197,7 +5237,9 @@ def _init_custom_logger_compatible_class( ) return NewRelicHandler.get_newrelic_logger_for_request( - standard_callback_dynamic_params=custom_logger_init_args, + standard_callback_dynamic_params=cast( # cast-ok: callback options come from dynamic config + StandardCallbackDynamicParams, custom_logger_init_args_value + ), in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache, ) @@ -5360,7 +5402,7 @@ def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list[CustomLogger]) def get_custom_logger_compatible_class( - logging_integration: _custom_logger_compatible_callbacks_literal, + logging_integration: str, ) -> CustomLogger | None: try: if logging_integration == "lago": @@ -5932,10 +5974,10 @@ class StandardLoggingPayloadSetup: elif isinstance(usage, Usage): return usage elif isinstance(usage, ResponseAPIUsage): - return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) elif isinstance(usage, dict): - if ResponseAPILoggingUtils._is_response_api_usage(usage): - return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + if ResponseAPILoggingUtils.is_response_api_usage(usage): + return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) if InteractionsUsageObjectTransformation.is_interactions_usage_object(usage): return InteractionsUsageObjectTransformation.transform_interactions_usage_object(usage) return Usage(**usage) @@ -5960,10 +6002,10 @@ class StandardLoggingPayloadSetup: if _raw is None: return _empty if isinstance(_raw, ResponseAPIUsage): - return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_raw).model_dump() + return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(_raw).model_dump() if isinstance(_raw, dict): - if ResponseAPILoggingUtils._is_response_api_usage(_raw): - return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_raw).model_dump() + if ResponseAPILoggingUtils.is_response_api_usage(_raw): + return ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(_raw).model_dump() if InteractionsUsageObjectTransformation.is_interactions_usage_object(_raw): return InteractionsUsageObjectTransformation.transform_interactions_usage_object(_raw).model_dump() return _raw @@ -5979,7 +6021,7 @@ class StandardLoggingPayloadSetup: init_response_obj: object, api_base: str | None = None, ) -> StandardLoggingModelInformation: - model_cost_name: Final = _select_model_name_for_cost_calc( + model_cost_name: Final = select_model_name_for_cost_calc( model=base_model if custom_pricing else None, completion_response=init_response_obj, base_model=base_model, @@ -6206,8 +6248,8 @@ class StandardLoggingPayloadSetup: error_code=error_status, error_class=error_class, llm_provider=_llm_provider_in_exception, - traceback=_redact_string(traceback_info), - error_message=_redact_string(error_message), + traceback=redact_string(traceback_info), + error_message=redact_string(error_message), error_provider_request_id=provider_request_id, error_rate_limit_category=rate_limit_category, error_rate_limit_type=rate_limit_type, @@ -6331,7 +6373,7 @@ class StandardLoggingPayloadSetup: return logging_obj.litellm_session_id @staticmethod - def _get_user_agent_tags(proxy_server_request: dict) -> list[str] | None: + def _get_user_agent_tags(proxy_server_request: Mapping[str, object]) -> list[str] | None: """ Return the user agent tags from the proxy server request for spend tracking """ @@ -6340,8 +6382,11 @@ class StandardLoggingPayloadSetup: user_agent_tags: list[str] | None = None headers: Final = proxy_server_request.get("headers", {}) if headers is not None and isinstance(headers, dict): - if "user-agent" in headers: - user_agent: Final = headers["user-agent"] + request_headers: Final = cast( # cast-ok: request headers are untyped framework data + dict[str, str | None], headers + ) + if "user-agent" in request_headers: + user_agent: Final = request_headers["user-agent"] if user_agent is not None: if user_agent_tags is None: user_agent_tags = [] @@ -6350,12 +6395,11 @@ class StandardLoggingPayloadSetup: user_agent_part = user_agent.split("/")[0] if user_agent_part is not None: user_agent_tags.append("User-Agent: " + user_agent_part) - if user_agent is not None: - user_agent_tags.append("User-Agent: " + user_agent) + user_agent_tags.append("User-Agent: " + user_agent) return user_agent_tags @staticmethod - def _get_extra_header_tags(proxy_server_request: dict) -> list[str] | None: + def _get_extra_header_tags(proxy_server_request: Mapping[str, object]) -> list[str] | None: """ Extract additional header tags for spend tracking based on config. """ @@ -6366,24 +6410,33 @@ class StandardLoggingPayloadSetup: headers: Final = proxy_server_request.get("headers", {}) if not isinstance(headers, dict): return None + request_headers: Final = cast( # cast-ok: request headers are untyped framework data + dict[str, str], headers + ) header_tags: Final = [] for header_name in extra_headers: - header_value = headers.get(header_name) + header_value = request_headers.get(header_name) if header_value: header_tags.append(f"{header_name}: {header_value}") return header_tags if header_tags else None @staticmethod - def _get_request_tags(litellm_params: dict, proxy_server_request: dict) -> list[str]: - # check for 'tags' in both 'metadata' and 'litellm_metadata' - metadata: Final = litellm_params.get("metadata") or {} - litellm_metadata: Final = litellm_params.get("litellm_metadata") or {} + def get_request_tags( + litellm_params: dict[str, object], + proxy_server_request: dict[str, object], + ) -> list[str]: + metadata: Final = cast( # cast-ok: request metadata is caller-provided + Mapping[str, object], litellm_params.get("metadata") or {} + ) + litellm_metadata: Final = cast( # cast-ok: request metadata is caller-provided + Mapping[str, object], litellm_params.get("litellm_metadata") or {} + ) if metadata.get("tags", []): - request_tags = metadata.get("tags", []).copy() + request_tags = cast(list[str], metadata.get("tags", [])).copy() # cast-ok: tags are caller-provided elif litellm_metadata.get("tags", []): - request_tags = litellm_metadata.get("tags", []).copy() + request_tags = cast(list[str], litellm_metadata.get("tags", [])).copy() # cast-ok: tags are caller-provided else: request_tags = [] user_agent_tags: Final = StandardLoggingPayloadSetup._get_user_agent_tags(proxy_server_request) @@ -6394,6 +6447,8 @@ class StandardLoggingPayloadSetup: request_tags.extend(additional_header_tags) return request_tags + _get_request_tags = get_request_tags + def _get_status_fields( status: StandardLoggingPayloadStatus, @@ -6575,7 +6630,7 @@ def get_standard_logging_object_payload( _model_id: Final = metadata.get("model_info", {}).get("id", "") _model_group: Final = metadata.get("model_group", "") - request_tags: Final = StandardLoggingPayloadSetup._get_request_tags( + request_tags: Final = StandardLoggingPayloadSetup.get_request_tags( litellm_params=litellm_params, proxy_server_request=proxy_server_request ) request_model_access_groups: Final = request_model_access_groups_from_litellm_params(litellm_params) @@ -6620,7 +6675,7 @@ def get_standard_logging_object_payload( if cache_hit is True: id = f"{id}_cache_hit{time.time()}" # do not duplicate the request id saved_cache_cost = ( - logging_obj._response_cost_calculator( + logging_obj.response_cost_calculator( result=init_response_obj, cache_hit=False, ) @@ -6628,7 +6683,7 @@ def get_standard_logging_object_payload( ) ## Get model cost information ## - base_model = _get_base_model_from_metadata(model_call_details=kwargs) + base_model = get_base_model_from_metadata(model_call_details=kwargs) # The router overrides completion_response.model to the model-group alias before # this payload is built, so cost-map lookup via that alias always misses. # Fall back to the actual deployment model set by the router in metadata. @@ -6681,7 +6736,9 @@ def get_standard_logging_object_payload( # Reconstruct full model name with provider prefix for logging # This ensures Bedrock models like "us.anthropic.claude-3-5-sonnet-20240620-v1:0" # are logged as "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0" - custom_llm_provider: Final = cast(str | None, kwargs.get("custom_llm_provider")) + custom_llm_provider: Final = cast( # cast-ok: provider name is caller-provided + str | None, kwargs.get("custom_llm_provider") + ) model_name = reconstruct_model_name(kwargs.get("model", "") or "", custom_llm_provider, metadata) response_model_name: str | None = None if isinstance(final_response_obj, dict): @@ -6933,7 +6990,7 @@ def _get_traceback_str_for_error(error_str: str) -> str: from decimal import Decimal # used for unit testing -from typing import Any, Union +from typing import Any def create_dummy_standard_logging_payload() -> StandardLoggingPayload: diff --git a/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py b/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py index bf99035a6b1..6c7760d7da8 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py +++ b/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py @@ -3,7 +3,7 @@ Helper utilities for tracking the cost of built-in tools. """ from collections.abc import Mapping -from typing import Final, Literal +from typing import Final, Literal, cast from pydantic import ValidationError @@ -739,9 +739,10 @@ class StandardBuiltInToolCostTracking: return False @staticmethod - def _get_web_search_options(kwargs: dict) -> WebSearchOptions | None: + def get_web_search_options(kwargs: Mapping[str, object]) -> WebSearchOptions | None: if "web_search_options" in kwargs: - return WebSearchOptions(**kwargs.get("web_search_options", {})) + web_search_options: Final = cast(WebSearchOptions, kwargs.get("web_search_options", {})) + return WebSearchOptions(**web_search_options) tools: Final = StandardBuiltInToolCostTracking._get_tools_from_kwargs( kwargs=kwargs, tool_type="web_search_preview" @@ -751,27 +752,31 @@ class StandardBuiltInToolCostTracking: for tool in tools: if isinstance(tool, dict): if StandardBuiltInToolCostTracking._is_web_search_tool_call(tool): - return WebSearchOptions(**tool) + return WebSearchOptions(**cast(WebSearchOptions, tool)) return None + _get_web_search_options = get_web_search_options + @staticmethod - def _get_tools_from_kwargs(kwargs: dict, tool_type: str) -> list[dict] | None: + def _get_tools_from_kwargs(kwargs: Mapping[str, object], tool_type: str) -> list[object] | None: if "tools" in kwargs: - return kwargs.get("tools", []) + return cast(list[object], kwargs.get("tools", [])) return None @staticmethod - def _get_file_search_tool_call(kwargs: dict) -> FileSearchTool | None: + def get_file_search_tool_call(kwargs: Mapping[str, object]) -> FileSearchTool | None: tools: Final = StandardBuiltInToolCostTracking._get_tools_from_kwargs(kwargs, "file_search") if tools: for tool in tools: if isinstance(tool, dict): if StandardBuiltInToolCostTracking._is_file_search_tool_call(tool): - return FileSearchTool(**tool) + return FileSearchTool(**cast(FileSearchTool, tool)) return None + _get_file_search_tool_call = get_file_search_tool_call + @staticmethod - def _is_web_search_tool_call(tool: dict) -> bool: + def _is_web_search_tool_call(tool: Mapping[str, object]) -> bool: if tool.get("type", None) == "web_search_preview": return True if tool.get("type", None) == "web_search": @@ -781,7 +786,7 @@ class StandardBuiltInToolCostTracking: return False @staticmethod - def _is_file_search_tool_call(tool: dict) -> bool: + def _is_file_search_tool_call(tool: Mapping[str, object]) -> bool: if tool.get("type", None) == "file_search": return True return False diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 29b5f1bd66a..368be9c3fd8 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -153,12 +153,15 @@ def get_web_search_requests_from_usage(usage: Usage) -> int | None: return get_web_search_requests(getattr(usage, "server_tool_use", None)) -def _is_above_128k(tokens: float) -> bool: +def is_above_128k(tokens: float) -> bool: if tokens > 128000: return True return False +_is_above_128k = is_above_128k + + def get_billable_input_tokens(usage: Usage) -> int: """ Returns the number of billable input tokens. @@ -185,7 +188,7 @@ def select_cost_metric_for_model( ) -def _generic_cost_per_character( +def generic_cost_per_character( model: str, custom_llm_provider: str, prompt_characters: float, @@ -250,7 +253,10 @@ def _generic_cost_per_character( return prompt_cost, completion_cost -def _get_service_tier_cost_key(base_key: str, service_tier: str | None) -> str: +_generic_cost_per_character = generic_cost_per_character + + +def get_service_tier_cost_key(base_key: str, service_tier: str | None) -> str: """ Get the appropriate cost key based on service tier. @@ -271,6 +277,9 @@ def _get_service_tier_cost_key(base_key: str, service_tier: str | None) -> str: return f"{base_key}_{suffix}" +_get_service_tier_cost_key = get_service_tier_cost_key + + def _parse_token_threshold(threshold: str) -> float: return float(threshold.replace("k", "")) * (1000 if "k" in threshold else 1) @@ -386,7 +395,7 @@ def _get_tiered_base_costs(model_info: ModelInfo, usage: Usage) -> tuple[float, completion_cost: Final = ( tier_rate(tier, "output_cost_per_token") if "output_cost_per_token" in tier - else _get_cost_per_unit(model_info, "output_cost_per_token") or 0.0 + else get_cost_per_unit(model_info, "output_cost_per_token") or 0.0 ) return ( tier_rate(tier, "input_cost_per_token"), @@ -644,30 +653,34 @@ def _get_token_base_cost( return _apply_off_peak_to_base_costs(model_info, current_time, tiered_base_costs) # Get service tier aware cost keys - input_cost_key: Final = _get_service_tier_cost_key("input_cost_per_token", service_tier) - output_cost_key: Final = _get_service_tier_cost_key("output_cost_per_token", service_tier) - cache_creation_cost_key: Final = _get_service_tier_cost_key("cache_creation_input_token_cost", service_tier) - cache_read_cost_key: Final = _get_service_tier_cost_key("cache_read_input_token_cost", service_tier) + input_cost_key: Final = get_service_tier_cost_key("input_cost_per_token", service_tier) + output_cost_key: Final = get_service_tier_cost_key("output_cost_per_token", service_tier) + cache_creation_cost_key: Final = get_service_tier_cost_key("cache_creation_input_token_cost", service_tier) + cache_read_cost_key: Final = get_service_tier_cost_key("cache_read_input_token_cost", service_tier) - prompt_base_cost = cast(float, _get_cost_per_unit(model_info, input_cost_key)) - completion_base_cost = cast(float, _get_cost_per_unit(model_info, output_cost_key)) + prompt_base_cost = cast( # cast-ok: model pricing data is external + float, get_cost_per_unit(model_info, input_cost_key) + ) + completion_base_cost = cast( # cast-ok: model pricing data is external + float, get_cost_per_unit(model_info, output_cost_key) + ) # For image generation models that don't have output_cost_per_token, # use output_cost_per_image_token as the base cost (all output tokens are image tokens) if completion_base_cost == 0.0 or completion_base_cost is None: - output_image_cost: Final = _get_cost_per_unit(model_info, "output_cost_per_image_token", None) + output_image_cost: Final = get_cost_per_unit(model_info, "output_cost_per_image_token", None) if output_image_cost is not None: - completion_base_cost = cast(float, output_image_cost) - cache_creation_cost = _get_cost_per_unit(model_info, cache_creation_cost_key, default_value=None) - cache_creation_cost_above_1hr = _get_cost_per_unit( + completion_base_cost = output_image_cost + cache_creation_cost = get_cost_per_unit(model_info, cache_creation_cost_key, default_value=None) + cache_creation_cost_above_1hr = get_cost_per_unit( model_info, "cache_creation_input_token_cost_above_1hr", default_value=None ) - cache_read_cost = _get_cost_per_unit(model_info, cache_read_cost_key, default_value=None) + cache_read_cost = get_cost_per_unit(model_info, cache_read_cost_key, default_value=None) ## CHECK IF ABOVE THRESHOLD # Optimization: collect threshold keys first to avoid sorting all model_info keys. # Standard thresholds and thresholds suffixed for this request's service tier both count. - tier_key_suffix: Final = _get_service_tier_cost_key("", service_tier) + tier_key_suffix: Final = get_service_tier_cost_key("", service_tier) threshold_keys: Final = [ k for k in model_info @@ -692,7 +705,7 @@ def _get_token_base_cost( # ON_DEMAND_PRIORITY. Falls back to the standard key automatically # via _get_cost_per_unit's service_tier fallback logic. tiered_input_key = ( - _get_service_tier_cost_key( + get_service_tier_cost_key( f"input_cost_per_token_above_{threshold_str}_tokens", service_tier, ) @@ -701,10 +714,10 @@ def _get_token_base_cost( ) prompt_base_cost = cast( float, - _get_cost_per_unit(model_info, tiered_input_key, prompt_base_cost), + get_cost_per_unit(model_info, tiered_input_key, prompt_base_cost), ) tiered_output_key = ( - _get_service_tier_cost_key( + get_service_tier_cost_key( f"output_cost_per_token_above_{threshold_str}_tokens", service_tier, ) @@ -713,7 +726,7 @@ def _get_token_base_cost( ) completion_base_cost = cast( float, - _get_cost_per_unit( + get_cost_per_unit( model_info, tiered_output_key, completion_base_cost, @@ -722,7 +735,7 @@ def _get_token_base_cost( # Apply tiered pricing to cache costs cache_creation_tiered_key = ( - _get_service_tier_cost_key( + get_service_tier_cost_key( f"cache_creation_input_token_cost_above_{threshold_str}_tokens", service_tier, ) @@ -730,7 +743,7 @@ def _get_token_base_cost( else f"cache_creation_input_token_cost_above_{threshold_str}_tokens" ) cache_creation_1hr_tiered_key = ( - _get_service_tier_cost_key( + get_service_tier_cost_key( f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens", service_tier, ) @@ -738,7 +751,7 @@ def _get_token_base_cost( else f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens" ) cache_read_tiered_key = ( - _get_service_tier_cost_key( + get_service_tier_cost_key( f"cache_read_input_token_cost_above_{threshold_str}_tokens", service_tier, ) @@ -746,13 +759,13 @@ def _get_token_base_cost( else f"cache_read_input_token_cost_above_{threshold_str}_tokens" ) - cache_creation_cost = _get_cost_per_unit(model_info, cache_creation_tiered_key, cache_creation_cost) + cache_creation_cost = get_cost_per_unit(model_info, cache_creation_tiered_key, cache_creation_cost) - cache_creation_cost_above_1hr = _get_cost_per_unit( + cache_creation_cost_above_1hr = get_cost_per_unit( model_info, cache_creation_1hr_tiered_key, cache_creation_cost_above_1hr ) - cache_read_cost = _get_cost_per_unit(model_info, cache_read_tiered_key, cache_read_cost) + cache_read_cost = get_cost_per_unit(model_info, cache_read_tiered_key, cache_read_cost) break except (IndexError, ValueError): @@ -795,13 +808,13 @@ def calculate_cost_component(model_info: ModelInfo, cost_key: str, usage_value: Returns: float: The calculated cost """ - cost_per_unit: Final = _get_cost_per_unit(model_info, cost_key) + cost_per_unit: Final = get_cost_per_unit(model_info, cost_key) if cost_per_unit is not None and isinstance(cost_per_unit, float) and usage_value is not None and usage_value > 0: return float(usage_value) * cost_per_unit return 0.0 -def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: float | None = 0.0) -> float | None: +def get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: float | None = 0.0) -> float | None: # Sometimes the cost per unit is a string (e.g.: If a value like "3e-7" was read from the config.yaml) cost_per_unit: Final = model_info.get(cost_key) if isinstance(cost_per_unit, float): @@ -834,7 +847,7 @@ def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: floa return float(fallback_cost) except ValueError: verbose_logger.exception( - "litellm.litellm_core_utils.llm_cost_calc.utils.py::_get_cost_per_unit(): Exception occured - %s\nDefaulting to 0.0", + "litellm.litellm_core_utils.llm_cost_calc.utils.py::get_cost_per_unit(): Exception occured - %s\nDefaulting to 0.0", fallback_cost, ) break # Only try the first matching suffix @@ -842,6 +855,9 @@ def _get_cost_per_unit(model_info: ModelInfo, cost_key: str, default_value: floa return default_value +_get_cost_per_unit = get_cost_per_unit + + def deployment_pricing(model_info: ModelInfo | None) -> ModelInfo | None: """The prices a deployment sets itself, as floats; None when it sets none that parse.""" if model_info is None: @@ -851,7 +867,7 @@ def deployment_pricing(model_info: ModelInfo | None) -> ModelInfo | None: { key: price for key in priced_keys - if (price := _get_cost_per_unit(model_info, key, default_value=None)) is not None + if (price := get_cost_per_unit(model_info, key, default_value=None)) is not None } ) if not pricing: @@ -868,7 +884,7 @@ def flat_image_cost(model_info: ModelInfo | None, image_response: ImageResponse) """The per-image price times the images returned; 0.0 when the table sets no per-image price.""" if model_info is None: return 0.0 - output_cost_per_image: Final = _get_cost_per_unit(model_info, "output_cost_per_image", default_value=None) or 0.0 + output_cost_per_image: Final = get_cost_per_unit(model_info, "output_cost_per_image", default_value=None) or 0.0 num_images: Final = len(image_response.data) if image_response.data else 0 return output_cost_per_image * num_images @@ -1086,9 +1102,9 @@ def _calculate_input_cost( ### CACHE READ COST - Now uses tiered pricing cache_hit_audio_tokens: Final = prompt_tokens_details["cache_hit_audio_tokens"] - audio_cache_read_rate: Final = _get_cost_per_unit( + audio_cache_read_rate: Final = get_cost_per_unit( model_info, - _get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier), + get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier), None, ) prompt_cost += float(prompt_tokens_details["cache_hit_tokens"] - cache_hit_audio_tokens) * cache_read_cost @@ -1100,7 +1116,7 @@ def _calculate_input_cost( if prompt_tokens_details["audio_tokens"] and not ( prompt_tokens_details["audio_length_seconds"] and model_info.get("input_cost_per_audio_per_second") is not None ): - audio_cost_key: Final = _get_service_tier_cost_key("input_cost_per_audio_token", service_tier) + audio_cost_key: Final = get_service_tier_cost_key("input_cost_per_audio_token", service_tier) prompt_cost += calculate_cost_component(model_info, audio_cost_key, prompt_tokens_details["audio_tokens"]) ### IMAGE TOKEN COST @@ -1173,7 +1189,7 @@ def _calculate_input_cost( return prompt_cost -def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str | None) -> float: +def get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str | None) -> float: """ Resolve the per-model regional-processing uplift multiplier for a given data-residency region. @@ -1195,7 +1211,7 @@ def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str | if multiplier is None: return 1.0 try: - return float(cast(float, multiplier)) + return float(multiplier) except (TypeError, ValueError): verbose_logger.exception( "Invalid regional_processing_uplift_multiplier_%s for model; defaulting to 1.0", @@ -1204,6 +1220,9 @@ def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str | return 1.0 +_get_regional_uplift_multiplier = get_regional_uplift_multiplier + + def get_vertex_regional_endpoint_uplift(model_info: ModelInfo, vertex_location: str | None) -> float: """ Resolve the per-model uplift multiplier for Vertex AI non-global (regional and @@ -1223,7 +1242,7 @@ def get_vertex_regional_endpoint_uplift(model_info: ModelInfo, vertex_location: if multiplier is None: return 1.0 try: - return float(cast(float, multiplier)) + return float(cast(float, multiplier)) # cast-ok: pricing multiplier is external model data except (TypeError, ValueError): verbose_logger.exception( "Invalid regional_endpoint_uplift_multiplier for model; defaulting to 1.0", @@ -1253,15 +1272,15 @@ def _resolve_reasoning_token_cost( service_tier: str | None, completion_base_cost: float, ) -> float: - tier_reasoning_key: Final = _get_service_tier_cost_key("output_cost_per_reasoning_token", service_tier) + tier_reasoning_key: Final = get_service_tier_cost_key("output_cost_per_reasoning_token", service_tier) if model_info.get(tier_reasoning_key) is not None: - tier_reasoning_cost: Final = _get_cost_per_unit(model_info, tier_reasoning_key, None) + tier_reasoning_cost: Final = get_cost_per_unit(model_info, tier_reasoning_key, None) if tier_reasoning_cost is not None: return tier_reasoning_cost - tier_output_key: Final = _get_service_tier_cost_key("output_cost_per_token", service_tier) + tier_output_key: Final = get_service_tier_cost_key("output_cost_per_token", service_tier) if tier_output_key != "output_cost_per_token" and model_info.get(tier_output_key) is not None: return completion_base_cost - standard_reasoning_cost: Final = _get_cost_per_unit(model_info, "output_cost_per_reasoning_token", None) + standard_reasoning_cost: Final = get_cost_per_unit(model_info, "output_cost_per_reasoning_token", None) return standard_reasoning_cost if standard_reasoning_cost is not None else completion_base_cost @@ -1444,7 +1463,7 @@ def generic_cost_per_token( ## AUDIO COST if not is_text_tokens_total and audio_tokens is not None and audio_tokens > 0: - _output_cost_per_audio_token = _get_cost_per_unit(resolved_model_info, "output_cost_per_audio_token", None) + _output_cost_per_audio_token = get_cost_per_unit(resolved_model_info, "output_cost_per_audio_token", None) _output_cost_per_audio_token = ( _output_cost_per_audio_token if _output_cost_per_audio_token is not None else completion_base_cost ) @@ -1462,7 +1481,7 @@ def generic_cost_per_token( ## IMAGE COST if not is_text_tokens_total and image_tokens and image_tokens > 0: - _output_cost_per_image_token = _get_cost_per_unit(resolved_model_info, "output_cost_per_image_token", None) + _output_cost_per_image_token = get_cost_per_unit(resolved_model_info, "output_cost_per_image_token", None) _output_cost_per_image_token = ( _output_cost_per_image_token if _output_cost_per_image_token is not None else completion_base_cost ) @@ -1470,7 +1489,7 @@ def generic_cost_per_token( ## VIDEO COST if not is_text_tokens_total and video_tokens and video_tokens > 0: - _output_cost_per_video_token = _get_cost_per_unit(resolved_model_info, "output_cost_per_video_token", None) + _output_cost_per_video_token = get_cost_per_unit(resolved_model_info, "output_cost_per_video_token", None) _output_cost_per_video_token = ( _output_cost_per_video_token if _output_cost_per_video_token is not None else completion_base_cost ) @@ -1479,7 +1498,7 @@ def generic_cost_per_token( ## REGIONAL DATA-RESIDENCY UPLIFT # Applied as a flat multiplier across all token costs for the request # when the upstream is a regionalized OpenAI host (eu./us.api.openai.com). - uplift: Final = _get_regional_uplift_multiplier(resolved_model_info, data_residency) + uplift: Final = get_regional_uplift_multiplier(resolved_model_info, data_residency) if uplift != 1.0: prompt_cost *= uplift completion_cost *= uplift @@ -1603,13 +1622,13 @@ def _cost_map_billed_rates( completion_base_cost=completion_base_cost, current_time=billing_time, ) - audio_cache_read_rate: Final = _get_cost_per_unit( + audio_cache_read_rate: Final = get_cost_per_unit( model_info, - _get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier), + get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier), None, ) multiplier: Final = ( - _get_regional_uplift_multiplier(model_info, data_residency) + get_regional_uplift_multiplier(model_info, data_residency) * get_vertex_regional_endpoint_uplift(model_info, vertex_location) * get_provider_specific_geo_multiplier(model_info=model_info, usage=usage) ) @@ -1764,7 +1783,7 @@ def calculate_prompt_caching_savings( cache_creation_cost_above_1hr=write_rate_1h - prompt_base_cost, cache_creation_cost=write_rate - prompt_base_cost, ) - uplift: Final = _get_regional_uplift_multiplier(model_info, data_residency) * get_vertex_regional_endpoint_uplift( + uplift: Final = get_regional_uplift_multiplier(model_info, data_residency) * get_vertex_regional_endpoint_uplift( model_info, vertex_location ) return (read_discount - write_premium) * uplift @@ -1899,7 +1918,7 @@ def calculate_image_response_web_search_cost( class CostCalculatorUtils: @staticmethod - def _call_type_has_image_response(call_type: str) -> bool: + def call_type_has_image_response(call_type: str) -> bool: """ Returns True if the call type has an image response @@ -1910,6 +1929,8 @@ class CostCalculatorUtils: """ return call_type in _IMAGE_RESPONSE_CALL_TYPES + _call_type_has_image_response = call_type_has_image_response + @staticmethod def route_image_generation_cost_calculator( model: str, diff --git a/litellm/litellm_core_utils/llm_request_utils.py b/litellm/litellm_core_utils/llm_request_utils.py index 7f9557003fd..0bf59a6d1f2 100644 --- a/litellm/litellm_core_utils/llm_request_utils.py +++ b/litellm/litellm_core_utils/llm_request_utils.py @@ -1,5 +1,5 @@ from collections.abc import Mapping -from typing import Final +from typing import Final, cast import litellm from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH @@ -106,7 +106,7 @@ def serialize_multipart_form_fields(data: Mapping[str, object]) -> tuple[tuple[s ) -def _ensure_extra_body_is_safe(extra_body: dict | None) -> dict | None: +def ensure_extra_body_is_safe(extra_body: dict[str, object] | None) -> dict[str, object] | None: """ Ensure that the extra_body sent in the request is safe, otherwise users will see this error @@ -117,22 +117,28 @@ def _ensure_extra_body_is_safe(extra_body: dict | None) -> dict | None: """ if extra_body is None: return None - if not isinstance(extra_body, dict): return extra_body - if "metadata" in extra_body and isinstance(extra_body["metadata"], dict): - if "prompt" in extra_body["metadata"]: - _prompt: Final = extra_body["metadata"].get("prompt") - + if "prompt" in cast(dict[str, object], extra_body["metadata"]): + prompt: Final = cast( # cast-ok: request metadata is caller-provided + dict[str, object], extra_body["metadata"] + ).get("prompt") # users can send Langfuse TextPromptClient objects, so we need to convert them to dicts # Langfuse TextPromptClients have .__dict__ attribute - if _prompt is not None and hasattr(_prompt, "__dict__"): - extra_body["metadata"]["prompt"] = _prompt.__dict__ + if prompt is not None and hasattr(prompt, "__dict__"): + cast(dict[str, object], extra_body["metadata"])["prompt"] = ( + cast( # cast-ok: prompt is an external SDK object + object, getattr(prompt, "__dict__") + ) + ) return extra_body +_ensure_extra_body_is_safe = ensure_extra_body_is_safe + + def pick_cheapest_chat_models_from_llm_provider(custom_llm_provider: str, n=1): """ Pick the n cheapest chat models from the LLM provider. diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 9ea730a873f..547329b7f51 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -10,7 +10,7 @@ import litellm from litellm._logging import verbose_logger from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _extract_reasoning_content, + extract_reasoning_content, ) from litellm.types.llms.databricks import DatabricksTool from litellm.types.llms.openai import ( @@ -71,7 +71,7 @@ def _normalize_images_for_message( return normalized -def _safe_convert_created_field(created_value) -> int: +def safe_convert_created_field(created_value: object) -> int: """ Safely convert a 'created' field value to an integer. @@ -91,19 +91,20 @@ def _safe_convert_created_field(created_value) -> int: elif isinstance(created_value, float): return int(created_value) else: - # for strings, etc try: - return int(float(created_value)) + return int(float(cast(float | str, created_value))) except (ValueError, TypeError): - # Fallback to current time if conversion fails return int(time.time()) +_safe_convert_created_field = safe_convert_created_field + + def convert_tool_call_to_json_mode( tool_calls: list[ChatCompletionMessageToolCall], convert_tool_call_to_json_mode: bool, ) -> tuple[Message | None, str | None]: - if _should_convert_tool_call_to_json_mode( + if should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=convert_tool_call_to_json_mode, ): @@ -245,7 +246,7 @@ async def convert_to_streaming_response_async( model_response_object.id = response_object["id"] if "created" in response_object: - model_response_object.created = _safe_convert_created_field(response_object["created"]) + model_response_object.created = safe_convert_created_field(response_object["created"]) if "system_fingerprint" in response_object: model_response_object.system_fingerprint = response_object["system_fingerprint"] @@ -334,7 +335,7 @@ def convert_to_streaming_response( model_response_object.id = response_object["id"] if "created" in response_object: - model_response_object.created = _safe_convert_created_field(response_object["created"]) + model_response_object.created = safe_convert_created_field(response_object["created"]) if "system_fingerprint" in response_object: model_response_object.system_fingerprint = response_object["system_fingerprint"] @@ -371,9 +372,9 @@ def convert_to_streaming_response( from collections import defaultdict -def _handle_invalid_parallel_tool_calls( +def handle_invalid_parallel_tool_calls( tool_calls: list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall], -): +) -> list[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | None: """ Handle hallucinated parallel tool call from openai - https://community.openai.com/t/model-tries-to-call-unknown-function-multi-tool-use-parallel/490653 @@ -414,6 +415,9 @@ def _handle_invalid_parallel_tool_calls( return tool_calls +_handle_invalid_parallel_tool_calls = handle_invalid_parallel_tool_calls + + class LiteLLMResponseObjectHandler: @staticmethod def convert_to_image_response( @@ -530,7 +534,7 @@ class LiteLLMResponseObjectHandler: return transformed_logprobs -def _should_convert_tool_call_to_json_mode( +def should_convert_tool_call_to_json_mode( tool_calls: ( Sequence[ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall] | Sequence[DatabricksTool] | None ) = None, @@ -545,6 +549,9 @@ def _should_convert_tool_call_to_json_mode( return False +_should_convert_tool_call_to_json_mode = should_convert_tool_call_to_json_mode + + def convert_to_model_response_object( response_object: dict | None = None, model_response_object: ModelResponse @@ -645,14 +652,14 @@ def convert_to_model_response_object( for _tc in tool_calls: _openai_tc = chat_completion_tool_call_from_dict(_tc) _openai_tool_calls.append(_openai_tc) - fixed_tool_calls = _handle_invalid_parallel_tool_calls(_openai_tool_calls) + fixed_tool_calls = handle_invalid_parallel_tool_calls(_openai_tool_calls) if fixed_tool_calls is not None: tool_calls = fixed_tool_calls message: Message | None = None finish_reason: str | None = None - if tool_calls is not None and _should_convert_tool_call_to_json_mode( + if tool_calls is not None and should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=convert_tool_call_to_json_mode, ): @@ -669,7 +676,7 @@ def convert_to_model_response_object( provider_specific_fields[f] = choice["message"][f] # Handle reasoning models that display `reasoning_content` within `content` - reasoning_content, content = _extract_reasoning_content(choice["message"]) + reasoning_content, content = extract_reasoning_content(choice["message"]) # Handle thinking models that display `thinking_blocks` within `content` thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = ( @@ -718,7 +725,7 @@ def convert_to_model_response_object( usage_object: Final = litellm.Usage(**response_object["usage"]) setattr(model_response_object, "usage", usage_object) if "created" in response_object: - model_response_object.created = _safe_convert_created_field(response_object["created"]) + model_response_object.created = safe_convert_created_field(response_object["created"]) if "id" in response_object: # Preserve the auto-generated id from ModelResponse.__init__ diff --git a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py index 7d9c33da923..b1749b5f679 100644 --- a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py +++ b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py @@ -125,7 +125,7 @@ class ResponseMetadata: "litellm_call_id": getattr(logging_obj, "litellm_call_id", None), "api_base": get_api_base(model=model or "", optional_params=kwargs), "model_id": model_id, - "response_cost": logging_obj._response_cost_calculator( + "response_cost": logging_obj.response_cost_calculator( result=self.result, litellm_model_name=model, router_model_id=model_id ), "additional_headers": process_response_headers( diff --git a/litellm/litellm_core_utils/logging_callback_manager.py b/litellm/litellm_core_utils/logging_callback_manager.py index 075dc83146f..5eaa4b467fe 100644 --- a/litellm/litellm_core_utils/logging_callback_manager.py +++ b/litellm/litellm_core_utils/logging_callback_manager.py @@ -38,7 +38,7 @@ class LoggingCallbackManager: except Exception: return False - def add_litellm_input_callback(self, callback: CustomLogger | str | Callable): + def add_litellm_input_callback(self, callback: CustomLogger | str | Callable[..., object]): """ Add a input callback to litellm.input_callback. Auto-routes async callbacks to litellm._async_input_callback. @@ -48,13 +48,13 @@ class LoggingCallbackManager: else: self._safe_add_callback_to_list(callback=callback, parent_list=litellm.input_callback) - def add_litellm_service_callback(self, callback: CustomLogger | str | Callable): + def add_litellm_service_callback(self, callback: CustomLogger | str | Callable[..., object]): """ Add a service callback to litellm.service_callback """ self._safe_add_callback_to_list(callback=callback, parent_list=litellm.service_callback) - def add_litellm_callback(self, callback: CustomLogger | str | Callable): + def add_litellm_callback(self, callback: CustomLogger | str | Callable[..., object]): """ Add a callback to litellm.callbacks @@ -65,7 +65,7 @@ class LoggingCallbackManager: parent_list=litellm.callbacks, ) - def add_litellm_success_callback(self, callback: CustomLogger | str | Callable): + def add_litellm_success_callback(self, callback: CustomLogger | str | Callable[..., object]): """ Add a success callback to `litellm.success_callback`. Auto-routes async callbacks to litellm._async_success_callback. @@ -81,7 +81,7 @@ class LoggingCallbackManager: else: self._safe_add_callback_to_list(callback=callback, parent_list=litellm.success_callback) - def add_litellm_failure_callback(self, callback: CustomLogger | str | Callable): + def add_litellm_failure_callback(self, callback: CustomLogger | str | Callable[..., object]): """ Add a failure callback to `litellm.failure_callback`. Auto-routes async callbacks to litellm._async_failure_callback. @@ -91,13 +91,13 @@ class LoggingCallbackManager: else: self._safe_add_callback_to_list(callback=callback, parent_list=litellm.failure_callback) - def add_litellm_async_success_callback(self, callback: CustomLogger | Callable | str): + def add_litellm_async_success_callback(self, callback: CustomLogger | Callable[..., object] | str): """ Add a success callback to litellm._async_success_callback """ self._safe_add_callback_to_list(callback=callback, parent_list=litellm._async_success_callback) - def add_litellm_async_failure_callback(self, callback: CustomLogger | Callable | str): + def add_litellm_async_failure_callback(self, callback: CustomLogger | Callable[..., object] | str): """ Add a failure callback to litellm._async_failure_callback """ @@ -139,7 +139,9 @@ class LoggingCallbackManager: for c in remove_list: callback_list.remove(c) - def _add_string_callback_to_list(self, callback: str, parent_list: list[CustomLogger | Callable | str]): + def _add_string_callback_to_list( + self, callback: str, parent_list: list[CustomLogger | Callable[..., object] | str] + ): """ Add a string callback to a list, if the callback is already in the list, do not add it again. """ @@ -148,7 +150,7 @@ class LoggingCallbackManager: else: verbose_logger.debug("Callback %s already exists in %s, not adding again..", callback, parent_list) - def _check_callback_list_size(self, parent_list: list[CustomLogger | Callable | str]) -> bool: + def _check_callback_list_size(self, parent_list: list[CustomLogger | Callable[..., object] | str]) -> bool: """ Check if adding another callback would exceed MAX_CALLBACKS Returns True if safe to add, False if would exceed limit @@ -163,7 +165,7 @@ class LoggingCallbackManager: return True @staticmethod - def _add_custom_callback_generic_api_str( + def add_custom_callback_generic_api_str( callback: str, ) -> GenericAPILogger | str: """ @@ -244,10 +246,12 @@ class LoggingCallbackManager: return callback + _add_custom_callback_generic_api_str = add_custom_callback_generic_api_str + def _safe_add_callback_to_list( self, - callback: CustomLogger | Callable | str, - parent_list: list[CustomLogger | Callable | str], + callback: CustomLogger | Callable[..., object] | str, + parent_list: list[CustomLogger | Callable[..., object] | str], ): """ Safe add a callback to a list, if the callback is already in the list, do not add it again. @@ -261,7 +265,7 @@ class LoggingCallbackManager: # Check if the callback is a custom callback if isinstance(callback, str): - callback = LoggingCallbackManager._add_custom_callback_generic_api_str(callback) + callback = LoggingCallbackManager.add_custom_callback_generic_api_str(callback) if isinstance(callback, str): self._add_string_callback_to_list(callback=callback, parent_list=parent_list) @@ -274,7 +278,11 @@ class LoggingCallbackManager: elif callable(callback): self._add_callback_function_to_list(callback=callback, parent_list=parent_list) - def _add_callback_function_to_list(self, callback: Callable, parent_list: list[CustomLogger | Callable | str]): + def _add_callback_function_to_list( + self, + callback: Callable[..., object], + parent_list: list[CustomLogger | Callable[..., object] | str], + ): """ Add a callback function to a list, if the callback is already in the list, do not add it again. """ @@ -289,7 +297,7 @@ class LoggingCallbackManager: def _add_custom_logger_to_list( self, custom_logger: CustomLogger, - parent_list: list[CustomLogger | Callable | str], + parent_list: list[CustomLogger | Callable[..., object] | str], ): """ Add a custom logger to a list, if another instance of the same custom logger exists in the list, do not add it again. @@ -341,7 +349,7 @@ class LoggingCallbackManager: litellm._async_failure_callback = [] litellm.callbacks = [] - def _get_all_callbacks(self) -> list[CustomLogger | Callable | str]: + def get_all_callbacks(self) -> list[CustomLogger | Callable[..., object] | str]: """ Get all callbacks from litellm.callbacks, litellm.success_callback, litellm.failure_callback, litellm._async_success_callback, litellm._async_failure_callback """ @@ -353,6 +361,8 @@ class LoggingCallbackManager: + litellm._async_failure_callback ) + _get_all_callbacks = get_all_callbacks + def remove_callback_from_all_lists(self, obj, require_self=False) -> None: """ Remove a callback object from every callback list it may have been @@ -379,7 +389,7 @@ class LoggingCallbackManager: Returns: Set[CustomLogger]: Set of custom loggers that are instances of the given class type """ - all_callbacks: Final = self._get_all_callbacks() + all_callbacks: Final = self.get_all_callbacks() matched_callbacks: Final[set[AdditionalLoggingUtils]] = set() for callback in all_callbacks: if isinstance(callback, CustomLogger) and isinstance(callback, AdditionalLoggingUtils): @@ -392,7 +402,7 @@ class LoggingCallbackManager: """ # ensure we don't have duplicate instances all_callbacks: Final = [] - for callback in self._get_all_callbacks(): + for callback in self.get_all_callbacks(): if isinstance(callback, callback_type) and callback not in all_callbacks: all_callbacks.append(callback) return all_callbacks @@ -401,7 +411,7 @@ class LoggingCallbackManager: """ Returns True if any of the active callbacks are of the given type """ - return any(isinstance(callback, callback_type) for callback in self._get_all_callbacks()) + return any(isinstance(callback, callback_type) for callback in self.get_all_callbacks()) def get_callbacks_by_type(self) -> CallbacksByType: """ @@ -444,11 +454,11 @@ class LoggingCallbackManager: def get_callback_objects(self) -> tuple[tuple[str, CustomLogger | Callable], ...]: return tuple( (self._get_callback_string(callback), callback) - for callback in self._get_all_callbacks() + for callback in self.get_all_callbacks() if not isinstance(callback, str) ) - def _get_callback_string(self, callback: CustomLogger | Callable | str) -> str: + def _get_callback_string(self, callback: CustomLogger | Callable[..., object] | str) -> str: from litellm.integrations.opentelemetry import OpenTelemetry from litellm.litellm_core_utils.custom_logger_registry import ( CustomLoggerRegistry, diff --git a/litellm/litellm_core_utils/logging_utils.py b/litellm/litellm_core_utils/logging_utils.py index 6aaa1ef692b..5e247324cea 100644 --- a/litellm/litellm_core_utils/logging_utils.py +++ b/litellm/litellm_core_utils/logging_utils.py @@ -5,7 +5,7 @@ import re import time from collections.abc import Iterator, Mapping, Sequence from datetime import datetime -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Any, Final, cast from litellm._logging import format_base64_size, verbose_logger from litellm.constants import ( @@ -199,10 +199,10 @@ def _get_parent_otel_span_from_logging_obj( # Reuse existing function by passing model_call_details as kwargs from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, + get_parent_otel_span_from_kwargs, ) - return _get_parent_otel_span_from_kwargs(logging_obj.model_call_details) + return get_parent_otel_span_from_kwargs(logging_obj.model_call_details) except Exception as e: verbose_logger.exception("Error in _get_parent_otel_span_from_logging_obj: %s", e) @@ -227,14 +227,14 @@ def convert_litellm_response_object_to_str( return None -def _assemble_complete_response_from_streaming_chunks( +def assemble_complete_response_from_streaming_chunks( result: ModelResponse | TextCompletionResponse | ModelResponseStream, start_time: datetime, end_time: datetime, - request_kwargs: dict, - streaming_chunks: list[Any], + request_kwargs: Mapping[str, object], + streaming_chunks: list[object], is_async: bool, -): +) -> ModelResponse | TextCompletionResponse | None: """ Assemble a complete response from a streaming chunks @@ -262,9 +262,10 @@ def _assemble_complete_response_from_streaming_chunks( if result.choices[0].finish_reason is not None: # if it's the last chunk streaming_chunks.append(result) try: + messages: Final = cast(list[dict[str, object]] | None, request_kwargs.get("messages", None)) complete_streaming_response = litellm.stream_chunk_builder( chunks=streaming_chunks, - messages=request_kwargs.get("messages", None), + messages=messages, start_time=start_time, end_time=end_time, ) @@ -279,6 +280,9 @@ def _assemble_complete_response_from_streaming_chunks( return complete_streaming_response +_assemble_complete_response_from_streaming_chunks = assemble_complete_response_from_streaming_chunks + + def _set_duration_in_model_call_details( logging_obj: Any, # we're not guaranteed this will be `LiteLLMLoggingObject` start_time: datetime, diff --git a/litellm/litellm_core_utils/model_param_helper.py b/litellm/litellm_core_utils/model_param_helper.py index 3696a328807..2dc52980875 100644 --- a/litellm/litellm_core_utils/model_param_helper.py +++ b/litellm/litellm_core_utils/model_param_helper.py @@ -1,5 +1,5 @@ from functools import lru_cache -from typing import Final +from typing import ClassVar, Final from openai.types.chat.completion_create_params import ( CompletionCreateParamsNonStreaming, @@ -24,7 +24,7 @@ from litellm.types.rerank import RerankRequest class ModelParamHelper: # Cached at class level — deterministic set built from static OpenAI type annotations - _relevant_logging_args: frozenset = frozenset() + relevant_logging_args: ClassVar[frozenset[str]] = frozenset() @staticmethod def get_standard_logging_model_parameters( @@ -32,7 +32,7 @@ class ModelParamHelper: ) -> dict: """ """ standard_logging_model_parameters: Final[dict] = {} - supported_model_parameters: Final = ModelParamHelper._relevant_logging_args + supported_model_parameters: Final = ModelParamHelper.relevant_logging_args for key, value in model_parameters.items(): if key in supported_model_parameters: @@ -44,20 +44,22 @@ class ModelParamHelper: return set(["messages", "prompt", "input", "system"]) @staticmethod - def _get_relevant_args_to_use_for_logging() -> set[str]: + def get_relevant_args_to_use_for_logging() -> set[str]: """ Gets all relevant llm api params besides the ones with prompt content """ - all_openai_llm_api_params: Final = ModelParamHelper._get_all_llm_api_params() + all_openai_llm_api_params: Final = ModelParamHelper.get_all_llm_api_params() # Exclude parameters that contain prompt content combined_kwargs: Final = all_openai_llm_api_params.difference( set(ModelParamHelper.get_exclude_params_for_model_parameters()) ) return combined_kwargs + _get_relevant_args_to_use_for_logging = get_relevant_args_to_use_for_logging + @staticmethod @lru_cache(maxsize=1) - def _get_all_llm_api_params() -> set[str]: + def get_all_llm_api_params() -> set[str]: """ Gets the supported kwargs for each call type and combines them. @@ -88,6 +90,8 @@ class ModelParamHelper: combined_kwargs = combined_kwargs.difference(exclude_kwargs) return combined_kwargs + _get_all_llm_api_params = get_all_llm_api_params + @staticmethod def get_litellm_provider_specific_params_for_chat_params() -> set[str]: return set(["thinking"]) @@ -185,4 +189,5 @@ class ModelParamHelper: return set(["metadata", "litellm_metadata"]) -ModelParamHelper._relevant_logging_args = frozenset(ModelParamHelper._get_relevant_args_to_use_for_logging()) +ModelParamHelper.relevant_logging_args = frozenset(ModelParamHelper.get_relevant_args_to_use_for_logging()) +ModelParamHelper._relevant_logging_args = ModelParamHelper.relevant_logging_args diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index d67b659fe37..49c4198c939 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -271,7 +271,7 @@ def request_contains_image_content(messages: Sequence[Mapping[str, object]]) -> ) -def _audio_or_image_in_message_content(message: AllMessageValues) -> bool: +def audio_or_image_in_message_content(message: AllMessageValues) -> bool: """ Checks if message content contains an image or audio """ @@ -284,6 +284,9 @@ def _audio_or_image_in_message_content(message: AllMessageValues) -> bool: return False +_audio_or_image_in_message_content = audio_or_image_in_message_content + + def convert_openai_message_to_only_content_messages( messages: list[AllMessageValues], ) -> list[dict[str, str]]: @@ -1503,7 +1506,7 @@ def tool_with_sanitized_parameters( return tool if sanitized_schema is input_schema else {**tool, "input_schema": sanitized_schema} -def _get_image_mime_type_from_url(url: str) -> str | None: +def get_image_mime_type_from_url(url: str) -> str | None: """ Get mime type for common image URLs See gemini mime types: https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/image-understanding#image-requirements @@ -1568,6 +1571,9 @@ def _get_image_mime_type_from_url(url: str) -> str | None: return None +_get_image_mime_type_from_url = get_image_mime_type_from_url + + def infer_content_type_from_url_and_content( url: str, content: bytes, @@ -1968,7 +1974,11 @@ def convert_prefix_message_to_non_prefix_messages( return new_messages -def _extract_reasoning_content(message: dict) -> tuple[str | None, str | None]: +def _provider_text_or_none(value: object) -> str | None: + return cast(str | None, value) # cast-ok: reasoning fields arrive in untyped provider messages + + +def extract_reasoning_content(message: Mapping[str, object]) -> tuple[str | None, str | None]: """ Extract reasoning content and main content from a message. @@ -1980,12 +1990,15 @@ def _extract_reasoning_content(message: dict) -> tuple[str | None, str | None]: """ message_content: Final = message.get("content") if "reasoning_content" in message: - return message["reasoning_content"], message_content + return _provider_text_or_none(message["reasoning_content"]), _provider_text_or_none(message_content) elif "reasoning" in message: - return message["reasoning"], message_content + return _provider_text_or_none(message["reasoning"]), _provider_text_or_none(message_content) elif isinstance(message_content, str): - return _parse_content_for_reasoning(message_content) - return None, message_content + return parse_content_for_reasoning(message_content) + return None, _provider_text_or_none(message_content) + + +_extract_reasoning_content = extract_reasoning_content def _readable_thinking_text(block: Mapping[str, object]) -> str: @@ -2163,7 +2176,7 @@ def responses_reasoning_items_from_thinking_blocks( ) -def _parse_content_for_reasoning( +def parse_content_for_reasoning( message_text: str | None, ) -> tuple[str | None, str | None]: """ @@ -2188,6 +2201,9 @@ def _parse_content_for_reasoning( return None, message_text +_parse_content_for_reasoning = parse_content_for_reasoning + + def _extract_base64_data(image_url: str) -> str: """ Extract pure base64 data from an image URL. diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 4f6d9af487b..f1eb9a0aef4 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -6,7 +6,7 @@ import json import mimetypes import re import xml.etree.ElementTree as ET -from collections.abc import Iterator, Mapping, Sequence +from collections.abc import Container, Iterator, Mapping, Sequence from enum import Enum from types import MappingProxyType from typing import Any, Final, TypeAlias, TypedDict, cast, overload @@ -459,7 +459,7 @@ async def _afetch_and_extract_template( Returns: (chat_template, bos_token, eos_token) """ from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import ( - _extract_token_value, + extract_token_value, ) bos_token = "" @@ -481,8 +481,8 @@ async def _afetch_and_extract_template( and "chat_template" in tokenizer_config["tokenizer"] ): tokenizer_data: dict = tokenizer_config["tokenizer"] - bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token")) - eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token")) + bos_token = extract_token_value(token_value=tokenizer_data.get("bos_token")) + eos_token = extract_token_value(token_value=tokenizer_data.get("eos_token")) chat_template = tokenizer_data["chat_template"] else: # Fallback: Try to fetch chat template from separate .jinja file @@ -496,8 +496,8 @@ async def _afetch_and_extract_template( and isinstance(tokenizer_config["tokenizer"], dict) ): tokenizer_data: dict = tokenizer_config["tokenizer"] - bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token")) - eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token")) + bos_token = extract_token_value(token_value=tokenizer_data.get("bos_token")) + eos_token = extract_token_value(token_value=tokenizer_data.get("eos_token")) else: raise Exception("No chat template found") @@ -513,7 +513,7 @@ def _fetch_and_extract_template( Returns: (chat_template, bos_token, eos_token) """ from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import ( - _extract_token_value, + extract_token_value, ) bos_token = "" @@ -535,8 +535,8 @@ def _fetch_and_extract_template( and "chat_template" in tokenizer_config["tokenizer"] ): tokenizer_data: dict = tokenizer_config["tokenizer"] - bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token")) - eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token")) + bos_token = extract_token_value(token_value=tokenizer_data.get("bos_token")) + eos_token = extract_token_value(token_value=tokenizer_data.get("eos_token")) chat_template = tokenizer_data["chat_template"] else: # Fallback: Try to fetch chat template from separate .jinja file @@ -550,8 +550,8 @@ def _fetch_and_extract_template( and isinstance(tokenizer_config["tokenizer"], dict) ): tokenizer_data: dict = tokenizer_config["tokenizer"] - bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token")) - eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token")) + bos_token = extract_token_value(token_value=tokenizer_data.get("bos_token")) + eos_token = extract_token_value(token_value=tokenizer_data.get("eos_token")) else: raise Exception("No chat template found") @@ -561,8 +561,8 @@ def _fetch_and_extract_template( async def ahf_chat_template(model: str, messages: list, chat_template: str | None = None): """HuggingFace chat template (async version)""" from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import ( - _aget_chat_template_file, - _aget_tokenizer_config, + aget_chat_template_file, + aget_tokenizer_config, strftime_now, ) @@ -573,8 +573,8 @@ async def ahf_chat_template(model: str, messages: list, chat_template: str | Non template, bos_token, eos_token = await _afetch_and_extract_template( model=model, chat_template=chat_template, - get_config_fn=_aget_tokenizer_config, - get_template_fn=_aget_chat_template_file, + get_config_fn=aget_tokenizer_config, + get_template_fn=aget_chat_template_file, ) return _render_chat_template( env=env, @@ -588,8 +588,8 @@ async def ahf_chat_template(model: str, messages: list, chat_template: str | Non def hf_chat_template(model: str, messages: list, chat_template: str | None = None): """HuggingFace chat template (sync version)""" from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import ( - _get_chat_template_file, - _get_tokenizer_config, + get_chat_template_file, + get_tokenizer_config, strftime_now, ) @@ -600,8 +600,8 @@ def hf_chat_template(model: str, messages: list, chat_template: str | None = Non template, bos_token, eos_token = _fetch_and_extract_template( model=model, chat_template=chat_template, - get_config_fn=_get_tokenizer_config, - get_template_fn=_get_chat_template_file, + get_config_fn=get_tokenizer_config, + get_template_fn=get_chat_template_file, ) return _render_chat_template( env=env, @@ -1161,7 +1161,7 @@ def _gemini_tool_call_invoke_helper( return function_call -def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: str | None) -> str: +def encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: str | None) -> str: """ Embed thought signature into tool call ID for OpenAI client compatibility. @@ -1180,7 +1180,10 @@ def _encode_tool_call_id_with_signature(tool_call_id: str, thought_signature: st return tool_call_id -def _get_thought_signature_from_tool(tool: dict) -> str | None: +_encode_tool_call_id_with_signature = encode_tool_call_id_with_signature + + +def get_thought_signature_from_tool(tool: Mapping[str, object]) -> str | None: """Extract thought signature from tool call's provider_specific_fields. If not provided try to extract thought signature from tool call id @@ -1192,26 +1195,41 @@ def _get_thought_signature_from_tool(tool: dict) -> str | None: # First check tool's provider_specific_fields provider_fields: Final = tool.get("provider_specific_fields") or {} if isinstance(provider_fields, dict): - signature = provider_fields.get("thought_signature") - if signature: - return signature + typed_provider_fields: Final = cast( # cast-ok: preserve dynamic provider response fields + dict[str, object], provider_fields + ) + signature_from_tool_fields: Final = typed_provider_fields.get("thought_signature") + if signature_from_tool_fields: + return cast(str, signature_from_tool_fields) # cast-ok: untyped provider response field # Then check function's provider_specific_fields function: Final = tool.get("function") if function: if isinstance(function, dict): - func_provider_fields: Final = function.get("provider_specific_fields") or {} + function_dict: Final = cast( # cast-ok: preserve dynamic provider response fields + dict[str, object], function + ) + func_provider_fields: Final = function_dict.get("provider_specific_fields") or {} if isinstance(func_provider_fields, dict): - signature = func_provider_fields.get("thought_signature") - if signature: - return signature - elif hasattr(function, "provider_specific_fields") and function.provider_specific_fields: - if isinstance(function.provider_specific_fields, dict): - signature = function.provider_specific_fields.get("thought_signature") - if signature: - return signature + typed_func_provider_fields: Final = cast( # cast-ok: preserve dynamic provider response fields + dict[str, object], func_provider_fields + ) + signature_from_function_fields: Final = typed_func_provider_fields.get("thought_signature") + if signature_from_function_fields: + return cast(str, signature_from_function_fields) # cast-ok: untyped provider response field + elif hasattr(function, "provider_specific_fields") and getattr(function, "provider_specific_fields"): + function_provider_fields: Final[object] = getattr(function, "provider_specific_fields") + if isinstance(function_provider_fields, dict): + typed_function_provider_fields: Final = cast( # cast-ok: provider fields are dynamic + dict[str, object], function_provider_fields + ) + signature_from_model_fields: Final = typed_function_provider_fields.get("thought_signature") + if signature_from_model_fields: + return cast(str, signature_from_model_fields) # cast-ok: untyped provider response field # Check if thought signature is embedded in tool call ID - tool_call_id: Final = tool.get("id") + tool_call_id: Final = cast( # cast-ok: tool IDs come from model responses + str, tool.get("id") + ) if tool_call_id and THOUGHT_SIGNATURE_SEPARATOR in tool_call_id: parts: Final = tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1) if len(parts) == 2: @@ -1220,6 +1238,9 @@ def _get_thought_signature_from_tool(tool: dict) -> str | None: return None +_get_thought_signature_from_tool = get_thought_signature_from_tool + + def _get_dummy_thought_signature() -> str: """Generate a dummy thought signature for models that require it. @@ -1301,7 +1322,7 @@ def convert_to_gemini_tool_call_invoke( ) if gemini_function_call is not None: part_dict: VertexPartType = {"function_call": gemini_function_call} - thought_signature = _get_thought_signature_from_tool(dict(tool)) + thought_signature = get_thought_signature_from_tool(dict(tool)) # Gemini signs only the first functionCall part of a parallel batch, so scope the # placeholder fallback to that part instead of fabricating one per sibling call: # https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/thinking/thought-signatures#parallel_function_calling_example @@ -1921,7 +1942,7 @@ def anthropic_infer_file_id_content_type( def anthropic_process_openai_file_message( message: ChatCompletionFileObject, ) -> AnthropicMessagesDocumentParam | AnthropicMessagesImageParam | AnthropicMessagesContainerUploadParam: - file_message: Final = cast(ChatCompletionFileObject, message) + file_message: Final = message file_sub: Final = file_message.get("file") if file_sub is None: raise litellm.BadRequestError( @@ -3383,7 +3404,7 @@ def _parse_content_type(content_type: str) -> str: return m.get_content_type() -def _parse_mime_type(base64_data: str) -> str | None: +def parse_mime_type(base64_data: str) -> str | None: mime_type_match: Final = re.match(r"data:(.*?);base64", base64_data) if mime_type_match: return mime_type_match.group(1) @@ -3391,6 +3412,9 @@ def _parse_mime_type(base64_data: str) -> str | None: return None +_parse_mime_type = parse_mime_type + + class BedrockImageProcessor: """Handles both sync and async image processing for Bedrock conversations.""" @@ -3461,7 +3485,7 @@ class BedrockImageProcessor: return img_without_base_64, mime_type, image_format @staticmethod - def _validate_format(mime_type: str, image_format: str) -> str: + def validate_format(mime_type: str, image_format: str) -> str: """Validate image format and mime type for both images and documents.""" supported_image_formats: Final = litellm.AmazonConverseConfig().get_supported_image_types() @@ -3488,6 +3512,8 @@ class BedrockImageProcessor: ) return image_format + _validate_format = validate_format + @staticmethod def _get_document_format(mime_type: str, supported_doc_formats: list[str]) -> str: """ @@ -3598,7 +3624,7 @@ class BedrockImageProcessor: mime_type = format image_format = mime_type.split("/")[1] - image_format = cls._validate_format(mime_type, image_format) + image_format = cls.validate_format(mime_type, image_format) return cls._create_bedrock_block(img_bytes, mime_type, image_format) @classmethod @@ -3617,7 +3643,7 @@ class BedrockImageProcessor: mime_type = format image_format = mime_type.split("/")[1] - image_format = cls._validate_format(mime_type, image_format) + image_format = cls.validate_format(mime_type, image_format) return cls._create_bedrock_block(img_bytes, mime_type, image_format) @@ -3889,7 +3915,7 @@ def _convert_to_bedrock_tool_call_result( tool_result: Final = BedrockToolResultBlock(content=tool_result_content_blocks, toolUseId=id) if used_search_results: - tool_result["status"] = cast(Literal["success"], "success") + tool_result["status"] = "success" content_block: Final = BedrockContentBlock(toolResult=tool_result) @@ -4082,7 +4108,7 @@ def _insert_assistant_continue_message( ) ) elif litellm.modify_params: - text = convert_content_list_to_str(cast(ChatCompletionAssistantMessage, DEFAULT_ASSISTANT_CONTINUE_MESSAGE)) + text = convert_content_list_to_str(DEFAULT_ASSISTANT_CONTINUE_MESSAGE) messages.append( BedrockMessageBlock( role="assistant", @@ -4407,14 +4433,14 @@ class BedrockConverseMessagesProcessor: _parts.append(_part) elif element["type"] == "file": _part = await BedrockConverseMessagesProcessor._async_process_file_message( - message=cast(ChatCompletionFileObject, element) + message=element ) _parts.append(_part) elif element["type"] == "document": _part = BedrockConverseMessagesProcessor._process_document_message(element) _parts.append(_part) _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( - message_block=cast(OpenAIMessageContentListBlock, element), + message_block=element, block_type="content_block", model=model, ) @@ -4528,7 +4554,7 @@ class BedrockConverseMessagesProcessor: if isinstance(element, dict): if element["type"] == "thinking": thinking_block = BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks( - thinking_blocks=[cast(ChatCompletionThinkingBlock, element)] + thinking_blocks=[element] ) assistants_parts = ( BedrockConverseMessagesProcessor.add_thinking_blocks_to_assistant_content( @@ -4611,7 +4637,7 @@ class BedrockConverseMessagesProcessor: return reasoning_content_blocks @staticmethod - def _process_file_message(message: ChatCompletionFileObject) -> BedrockContentBlock: + def process_file_message(message: ChatCompletionFileObject) -> BedrockContentBlock: file_message: Final = message.get("file") if file_message is None: raise litellm.BadRequestError( @@ -4631,6 +4657,8 @@ class BedrockConverseMessagesProcessor: format: Final = file_message.get("format") return BedrockImageProcessor.process_image_sync(image_url=cast(str, file_id or file_data), format=format) + _process_file_message = process_file_message + @staticmethod async def _async_process_file_message( message: ChatCompletionFileObject, @@ -4669,7 +4697,7 @@ class BedrockConverseMessagesProcessor: ) media_type: Final[str] = source["media_type"] data: Final[str] = source["data"] - doc_format = BedrockImageProcessor._validate_format(mime_type=media_type, image_format=media_type.split("/")[1]) + doc_format = BedrockImageProcessor.validate_format(mime_type=media_type, image_format=media_type.split("/")[1]) # Deterministic name using the same hashing pattern as _create_bedrock_block HASH_SAMPLE_BYTES: Final = 64 * 1024 @@ -4780,15 +4808,13 @@ def _bedrock_converse_messages_pt( ) _parts.append(_part) elif element["type"] == "file": - _part = BedrockConverseMessagesProcessor._process_file_message( - message=cast(ChatCompletionFileObject, element) - ) + _part = BedrockConverseMessagesProcessor.process_file_message(message=element) _parts.append(_part) elif element["type"] == "document": _part = BedrockConverseMessagesProcessor._process_document_message(element) _parts.append(_part) _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( - message_block=cast(OpenAIMessageContentListBlock, element), + message_block=element, block_type="content_block", model=model, ) @@ -4905,7 +4931,7 @@ def _bedrock_converse_messages_pt( if element["type"] == "thinking": thinking_block = ( BedrockConverseMessagesProcessor.translate_thinking_blocks_to_reasoning_content_blocks( - thinking_blocks=[cast(ChatCompletionThinkingBlock, element)] + thinking_blocks=[element] ) ) assistants_parts = ( @@ -5166,18 +5192,33 @@ def _bedrock_tools_pt(tools: list, model: str | None = None) -> list[BedrockTool # Function call template -def function_call_prompt(messages: list, functions: list): +def function_call_prompt( + messages: list[dict[str, object]], + functions: list[object], +) -> list[dict[str, object]]: function_prompt = """Produce JSON OUTPUT ONLY! Adhere to this format {"name": "function_name", "arguments":{"argument_name": "argument_value"}} The following functions are available to you:""" for function in functions: function_prompt += f"""\n{function}\n""" + def _append_function_prompt(message: dict[str, object]) -> bool: + role: Final = cast( # cast-ok: preserve dynamic role membership behavior + Container[object], message["role"] + ) + if "system" not in role: + return False + + content: Final = message["content"] + if isinstance(content, str): + message["content"] = f"{content} {function_prompt}" + else: + cast( # cast-ok: preserve dynamic content append behavior + list[object], content + ).append({"type": "text", "text": f""" {function_prompt}"""}) + return True + function_added_to_prompt = False for message in messages: - if "system" in message["role"]: - if isinstance(message["content"], str): - message["content"] += f""" {function_prompt}""" - else: - message["content"].append({"type": "text", "text": f""" {function_prompt}"""}) + if _append_function_prompt(message): function_added_to_prompt = True if function_added_to_prompt is False: diff --git a/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py b/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py index 8f8228d6dfd..1f54f62392d 100644 --- a/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py +++ b/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py @@ -38,7 +38,7 @@ def strftime_now(fmt: str) -> str: return datetime.now().strftime(fmt) -def _get_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult: +def get_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult: """ Fetch tokenizer_config.json from HuggingFace (sync) @@ -61,7 +61,10 @@ def _get_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult: return {"status": "failure"} -async def _aget_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult: +_get_tokenizer_config = get_tokenizer_config + + +async def aget_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult: """ Fetch tokenizer_config.json from HuggingFace (async) @@ -86,7 +89,10 @@ async def _aget_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult: return {"status": "failure"} -def _get_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult: +_aget_tokenizer_config = aget_tokenizer_config + + +def get_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult: """ Fetch chat template from separate .jinja file (sync) @@ -114,7 +120,10 @@ def _get_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult: return {"status": "failure"} -async def _aget_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult: +_get_chat_template_file = get_chat_template_file + + +async def aget_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult: """ Fetch chat template from separate .jinja file (async) @@ -144,7 +153,10 @@ async def _aget_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResul return {"status": "failure"} -def _extract_token_value(token_value: None | str | dict[str, Any]) -> str: +_aget_chat_template_file = aget_chat_template_file + + +def extract_token_value(token_value: None | str | dict[str, Any]) -> str: """ Extract token string from various formats (string, dict, etc.) @@ -159,3 +171,6 @@ def _extract_token_value(token_value: None | str | dict[str, Any]) -> str: if isinstance(token_value, dict): return token_value.get("content", "") return "" + + +_extract_token_value = extract_token_value diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index d2fbb26bb02..0f5551fe446 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -150,7 +150,7 @@ class RealTimeStreaming: self.tool_calls: list[dict] = [] # Detect whether the client is explicitly opting into the beta protocol. - self._client_wants_beta = self._detect_beta_header(websocket) + self._client_wants_beta = self.detect_beta_header(websocket) self._backend_uses_beta_protocol = ( self._client_wants_beta if backend_uses_beta_protocol is None else backend_uses_beta_protocol ) @@ -178,7 +178,7 @@ class RealTimeStreaming: self._pending_guardrail_message: str | None = None # Track whether session.created has already been sent to the client # (e.g. synthetic event in deferred setup mode). - self._session_created_sent_to_client: bool = False + self.session_created_sent_to_client: bool = False # Track whether we have already sent the guardrail turn-detection update # that disables provider auto-response for transcription guardrails. self._guardrail_turn_detection_update_sent: bool = False @@ -199,6 +199,14 @@ class RealTimeStreaming: # Optional per-provider GA event normalizer (e.g. XAIRealtimeNormalizer). self._event_normalizer = event_normalizer + @property + def _session_created_sent_to_client(self) -> bool: + return self.session_created_sent_to_client + + @_session_created_sent_to_client.setter + def _session_created_sent_to_client(self, value: bool) -> None: + self.session_created_sent_to_client = value + # Per-connection caps for pre-setup audio frames (message count + total bytes). _MAX_BUFFERED_MESSAGES: int = 200 _MAX_BUFFERED_BYTES: int = 10 * 1024 * 1024 # 10 MB @@ -252,7 +260,7 @@ class RealTimeStreaming: # TypedDict union members do not narrow to plain dict for mypy. message_obj: dict[str, Any] = cast(dict[str, Any], message) else: - message_obj = cast(dict[str, Any], json.loads(cast(str, message))) + message_obj = cast(dict[str, Any], json.loads(message)) self._collect_tool_calls_from_response_done(cast(dict, message_obj)) if not self._should_store_message(message_obj): return @@ -457,7 +465,7 @@ class RealTimeStreaming: if self._content_sent_after_setup: verbose_logger.debug("Dropping follow-up setup after content was already sent to backend") continue - msg = self._maybe_inject_guardrail_auto_response_disable(msg) + msg = self.maybe_inject_guardrail_auto_response_disable(msg) await self.backend_ws.send(msg) self._cache_session_configuration_request(msg) sent = True @@ -734,7 +742,7 @@ class RealTimeStreaming: if sent: self._guardrail_turn_detection_update_sent = True - def _maybe_inject_guardrail_auto_response_disable(self, setup_message: str) -> str: + def maybe_inject_guardrail_auto_response_disable(self, setup_message: str) -> str: """Fold the transcription-guardrail auto-response disable into the setup. Gemini/Vertex Live reject a second ``setup`` (1007), so the guardrail's @@ -764,6 +772,8 @@ class RealTimeStreaming: ) return json.dumps(obj) + _maybe_inject_guardrail_auto_response_disable = maybe_inject_guardrail_auto_response_disable + def _has_realtime_guardrails_for_event_hooks( self, event_hooks: Sequence["GuardrailEventHooks"], @@ -981,7 +991,7 @@ class RealTimeStreaming: break finally: self._flushing_pending_messages_until_setup = False - if self._session_created_sent_to_client: + if self.session_created_sent_to_client: # A synthetic session.created (with placeholder defaults) was # already forwarded to the client when we connected. The # provider's real session.created (e.g. emitted from Gemini @@ -991,7 +1001,7 @@ class RealTimeStreaming: # configuration without seeing two `session.created` events. event = {**event, "type": "session.updated"} else: - self._session_created_sent_to_client = True + self.session_created_sent_to_client = True event_str = json.dumps(event) ## For audio/VAD guardrail path: forward the (possibly retyped) ## session.created first, then invoke the one-time guardrail @@ -1154,7 +1164,7 @@ class RealTimeStreaming: self.logging_obj.model_call_details[REALTIME_SESSION_FAILURE_LOGGED_KEY] = True @staticmethod - def _detect_beta_header(websocket: ScopedWebSocket) -> bool: + def detect_beta_header(websocket: ScopedWebSocket) -> bool: """Return True if the client sent 'OpenAI-Beta: realtime=v1'. Checks the raw ASGI scope headers so it works for both FastAPI WebSocket @@ -1173,6 +1183,8 @@ class RealTimeStreaming: pass return False + _detect_beta_header = detect_beta_header + @staticmethod def _remap_beta_session_to_ga(session: dict) -> dict: """ @@ -1591,4 +1603,4 @@ class RealTimeStreaming: def client_sent_openai_beta_realtime_header(websocket: ScopedWebSocket) -> bool: """True when the client WebSocket includes ``OpenAI-Beta: realtime=v1``.""" - return RealTimeStreaming._detect_beta_header(websocket) + return RealTimeStreaming.detect_beta_header(websocket) diff --git a/litellm/litellm_core_utils/secret_redaction.py b/litellm/litellm_core_utils/secret_redaction.py index 89df5e500db..2b782f0a4b9 100644 --- a/litellm/litellm_core_utils/secret_redaction.py +++ b/litellm/litellm_core_utils/secret_redaction.py @@ -88,13 +88,16 @@ def _build_secret_patterns() -> "re.Pattern[str]": _SECRET_RE: Final = _build_secret_patterns() -def _python_redact_string(value: str) -> str: +def python_redact_string(value: str) -> str: return _SECRET_RE.sub(REDACTED, value) +_python_redact_string = python_redact_string + + def redact_string(value: str) -> str: """Scrub known secret/credential patterns from *value* and return the result.""" - return diagnostics.run(lambda native: native.redact_text(value), lambda: _python_redact_string(value)) + return diagnostics.run(lambda native: native.redact_text(value), lambda: python_redact_string(value)) _UNIX_SYSTEM_PATH: Final = r"/(?:etc|var|opt|usr|home|root|private|Users|tmp|mnt|srv)/[^\s'\"\)\]}>,]+" @@ -115,7 +118,7 @@ def _python_redact_internal_details(value: str) -> str: on top of redact_string(). For client-facing messages only: server logs keep this detail.""" marker_index: Final = value.find(_TRACEBACK_MARKER) without_traceback: Final = value[:marker_index].rstrip() if marker_index != -1 else value - return _INTERNAL_DETAIL_RE.sub(REDACTED, _python_redact_string(without_traceback)) + return _INTERNAL_DETAIL_RE.sub(REDACTED, python_redact_string(without_traceback)) def redact_internal_details(value: str) -> str: @@ -124,7 +127,7 @@ def redact_internal_details(value: str) -> str: ) -def _python_redact_structured_value(key: str | None, value: str) -> str: +def python_redact_structured_value(key: str | None, value: str) -> str: """Scrub *value* as it appeared under *key* inside a structured record. redact_string() replaces a whole ``key: value`` span with REDACTED, which is @@ -133,15 +136,18 @@ def _python_redact_structured_value(key: str | None, value: str) -> str: repr would, so the key-name patterns still fire, but collapses only the value so the caller's structure survives. """ - scrubbed: Final = _python_redact_string(value) + scrubbed: Final = python_redact_string(value) if scrubbed != value or key is None: return scrubbed rendered: Final = f"'{key}': '{value}'" - return REDACTED if _python_redact_string(rendered) != rendered else value + return REDACTED if python_redact_string(rendered) != rendered else value + + +_python_redact_structured_value = python_redact_structured_value def redact_structured_value(key: str | None, value: str) -> str: return diagnostics.run( lambda native: native.redact_structured_text(key, value), - lambda: _python_redact_structured_value(key, value), + lambda: python_redact_structured_value(key, value), ) diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py index 828fd09d3ad..ef18cda0c6c 100644 --- a/litellm/litellm_core_utils/sensitive_data_masker.py +++ b/litellm/litellm_core_utils/sensitive_data_masker.py @@ -54,7 +54,7 @@ class SensitiveDataMasker: self.mask_char = mask_char self.mask_short_values = mask_short_values - def _mask_value(self, value: str) -> str: + def mask_value(self, value: str) -> str: value_str: Final = str(value) if not value_str: return value @@ -71,6 +71,8 @@ class SensitiveDataMasker: f"{value_str[: self.visible_prefix]}{self.mask_char * masked_length}{value_str[-self.visible_suffix :]}" ) + _mask_value = mask_value + def is_sensitive_key(self, key: str, excluded_keys: set[str] | None = None) -> bool: # Check if key is in excluded_keys first (exact match) if excluded_keys and key in excluded_keys: @@ -109,7 +111,7 @@ class SensitiveDataMasker: elif isinstance(item, list): masked_items.append(self._mask_sequence(item, depth + 1, max_depth, excluded_keys, key_is_sensitive)) elif key_is_sensitive and isinstance(item, str): - masked_items.append(self._mask_value(item)) + masked_items.append(self.mask_value(item)) else: masked_items.append(item if isinstance(item, (int, float, bool, str, list)) else str(item)) return masked_items @@ -136,7 +138,7 @@ class SensitiveDataMasker: masked_data[k] = self.mask_dict(vars(v), depth + 1, max_depth, excluded_keys) elif key_is_sensitive: str_value = str(v) if v is not None else "" - masked_data[k] = self._mask_value(str_value) + masked_data[k] = self.mask_value(str_value) else: masked_data[k] = v if isinstance(v, (int, float, bool, str, list)) else str(v) except Exception: @@ -198,7 +200,7 @@ class _PayloadWalker: def walk(self, node: object, key_is_sensitive: bool, depth: int) -> object: if not isinstance(node, (Mapping, list, tuple, BaseModel)): - return _default_masker._mask_value(node) if key_is_sensitive and isinstance(node, str) and node else node + return _default_masker.mask_value(node) if key_is_sensitive and isinstance(node, str) and node else node if depth >= DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER: return REDACTED memo_key: Final = (id(node), key_is_sensitive and not isinstance(node, Mapping)) @@ -242,7 +244,7 @@ def mask_sensitive_keys(data: Mapping[str, object], sensitive_fields: set[str]) if len(value) < min_visible: masked[key] = mask_char * len(value) if value else value else: - masked[key] = _default_masker._mask_value(value) + masked[key] = _default_masker.mask_value(value) else: masked[key] = value return masked diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 6080d52ea89..28402624297 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -698,7 +698,7 @@ class CustomStreamWrapper: if isinstance(chunk, bytes): chunk = chunk.decode("utf-8") if "text_output" in chunk: - response = CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or "" + response = CustomStreamWrapper.strip_sse_data_from_chunk(chunk) or "" response = response.strip() parsed_response = json.loads(response) else: @@ -1676,15 +1676,15 @@ class CustomStreamWrapper: """ Caches the streaming response """ - if not cache_hit and self.logging_obj._llm_caching_handler is not None: - self.logging_obj._llm_caching_handler._sync_add_streaming_response_to_cache(processed_chunk) + if not cache_hit and self.logging_obj.llm_caching_handler is not None: + self.logging_obj.llm_caching_handler.sync_add_streaming_response_to_cache(processed_chunk) async def async_cache_streaming_response(self, processed_chunk, cache_hit: bool): """ Caches the streaming response """ - if not cache_hit and self.logging_obj._llm_caching_handler is not None: - await self.logging_obj._llm_caching_handler._add_streaming_response_to_cache(processed_chunk) + if not cache_hit and self.logging_obj.llm_caching_handler is not None: + await self.logging_obj.llm_caching_handler.add_streaming_response_to_cache(processed_chunk) def run_success_logging_and_cache_storage(self, processed_chunk, cache_hit: bool): """ @@ -1710,7 +1710,7 @@ class CustomStreamWrapper: asyncio.run(self.logging_obj.async_success_handler(processed_chunk, None, None, cache_hit)) ## SYNC LOGGING — only for sync SDK entrypoints; async proxy paths export via async_success_handler litellm_params: Final = self.logging_obj.model_call_details.get("litellm_params", {}) - if self.logging_obj._is_sync_litellm_request(litellm_params): + if self.logging_obj.is_sync_litellm_request(litellm_params): self.logging_obj.success_handler(processed_chunk, None, None, cache_hit) def finish_reason_handler(self): @@ -1794,7 +1794,7 @@ class CustomStreamWrapper: if response is None: continue if self.logging_obj.completion_start_time is None: - self.logging_obj._update_completion_start_time(completion_start_time=datetime.datetime.now()) + self.logging_obj.update_completion_start_time(completion_start_time=datetime.datetime.now()) ## LOGGING if not litellm.disable_streaming_logging: executor.submit( @@ -2005,7 +2005,7 @@ class CustomStreamWrapper: continue if self.logging_obj.completion_start_time is None: - self.logging_obj._update_completion_start_time(completion_start_time=datetime.datetime.now()) + self.logging_obj.update_completion_start_time(completion_start_time=datetime.datetime.now()) if processed_chunk.choices: choice = processed_chunk.choices[0] @@ -2252,7 +2252,7 @@ class CustomStreamWrapper: backfill_missing_cache_usage_fields(usage) self.logging_obj.model_call_details["combined_usage_object"] = usage self.logging_obj.model_call_details["response_cost"] = ( - self.logging_obj._response_cost_calculator(result=partial_response) or 0.0 + self.logging_obj.response_cost_calculator(result=partial_response) or 0.0 ) except Exception as recover_error: verbose_logger.debug( @@ -2334,7 +2334,7 @@ class CustomStreamWrapper: ) @staticmethod - def _strip_sse_data_from_chunk(chunk: str | None) -> str | None: + def strip_sse_data_from_chunk(chunk: str | None) -> str | None: """ Strips the 'data: ' prefix from Server-Sent Events (SSE) chunks. @@ -2369,6 +2369,8 @@ class CustomStreamWrapper: return chunk + _strip_sse_data_from_chunk = strip_sse_data_from_chunk + def _cache_token_count(details: PromptTokensDetailsWrapper | None, keys: tuple[str, ...]) -> int: for key in keys: diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index b92eb74cbb1..7357eb5ef19 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -15,7 +15,7 @@ from typing_extensions import ParamSpec, TypeVar import litellm from litellm import verbose_logger -from litellm._lazy_imports import _get_default_encoding +from litellm._lazy_imports import get_default_encoding from litellm.constants import ( DEFAULT_IMAGE_HEIGHT, DEFAULT_IMAGE_TOKEN_COUNT, @@ -650,35 +650,34 @@ def _get_exact_count_function( ) -> TokenCounterFunction: """ Get the function to count tokens based on the model and custom tokenizer.""" - from litellm.utils import _select_tokenizer + from litellm.utils import select_tokenizer if model is not None or custom_tokenizer is not None: - tokenizer_json: Final = custom_tokenizer or _select_tokenizer(model) - if tokenizer_json["type"] == "huggingface_tokenizer": - tokenizer: Final[HuggingFace] = tokenizer_json["tokenizer"] - - def count_tokens(text: str) -> int: - if isinstance(tokenizer, HuggingFaceTokenizer): - return tokenizer.count(text) - return len(tokenizer.encode_batch_fast([text])[0]) - - return count_tokens - elif tokenizer_json["type"] == "openai_tokenizer": - encoding: Final = openai_tokenizer_encoding(model) - - def encode_length(text: str) -> int: - return _encoding_count(encoding, text) - - return _get_tiktoken_count_function(encode_length) - else: - raise ValueError("Unsupported tokenizer type") + tokenizer_json: Final = custom_tokenizer or select_tokenizer(model) else: - default_encoding: Final = _get_default_encoding() + default_encoding: Final = get_default_encoding() def encode_length(text: str) -> int: return _encoding_count(default_encoding, text) return _get_tiktoken_count_function(encode_length) + if tokenizer_json["type"] == "huggingface_tokenizer": + tokenizer: Final[HuggingFace] = tokenizer_json["tokenizer"] + + def count_tokens(text: str) -> int: + if isinstance(tokenizer, HuggingFaceTokenizer): + return tokenizer.count(text) + return len(tokenizer.encode_batch_fast([text])[0]) + + return count_tokens + if tokenizer_json["type"] == "openai_tokenizer": + encoding: Final = openai_tokenizer_encoding(model) + + def encode_length(text: str) -> int: + return _encoding_count(encoding, text) + + return _get_tiktoken_count_function(encode_length) + raise ValueError("Unsupported tokenizer type") def _encoding_count(encoding: Encoding, text: str) -> int: diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 53e8605011a..24c8bc07c76 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -52,7 +52,7 @@ from litellm.types.utils import ( ModelResponseStream, StreamingChoices, Usage, - _generate_id, + generate_id, ) from ...base import BaseLLM @@ -566,7 +566,7 @@ class ModelResponseIterator: # common case (no '/' or other invalid chars in any tool name). self.tool_name_reverse_map: dict[str, str] = tool_name_reverse_map or {} # Generate response ID once per stream to match OpenAI-compatible behavior - self.response_id = _generate_id() + self.response_id = generate_id() self.served_model: str | None = None # Track if we're currently streaming a response_format tool diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 3b46fa9b9df..69016f5c211 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -754,11 +754,11 @@ class AnthropicModelInfo(BaseLLMModelInfo): def _get_model_capability(model: str, key: str) -> bool | None: """Read boolean capability ``key`` from the model map, or None when no entry declares it.""" - from litellm.utils import _get_bundled_model_cost_map + from litellm.utils import get_bundled_model_cost_map try: candidates: Final = AnthropicModelInfo._model_map_lookup_candidates(model) - for model_cost in (litellm.model_cost, _get_bundled_model_cost_map()): + for model_cost in (litellm.model_cost, get_bundled_model_cost_map()): for cand in candidates: value = model_cost.get(cand, {}).get(key) if isinstance(value, bool): @@ -793,13 +793,13 @@ class AnthropicModelInfo(BaseLLMModelInfo): model does not resolve under that provider or the resolved entry has no opinion on ``key``. """ - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper try: resolved_model, resolved_provider, _, _ = litellm.get_llm_provider( model=model, custom_llm_provider=custom_llm_provider ) - value: Final = _get_model_info_helper(model=resolved_model, custom_llm_provider=resolved_provider).get(key) + value: Final = get_model_info_helper(model=resolved_model, custom_llm_provider=resolved_provider).get(key) except Exception: # noqa: BLE001 # _get_model_info_helper raises bare Exception for unmapped models return None return value if isinstance(value, bool) else None @@ -813,13 +813,13 @@ class AnthropicModelInfo(BaseLLMModelInfo): Otherwise ``_supports_factory``'s provider-level fallbacks and the raw model-map walk remain as backstops for alias forms the lookup misses. """ - from litellm.utils import _supports_factory + from litellm.utils import supports_factory resolved: Final = AnthropicModelInfo._get_provider_resolved_capability(model, key, custom_llm_provider) if resolved is not None: return resolved try: - if _supports_factory( + if supports_factory( model=model, custom_llm_provider=custom_llm_provider, key=key, diff --git a/litellm/llms/anthropic/pass_through/messages/handler.py b/litellm/llms/anthropic/pass_through/messages/handler.py index 3068a7eb3ab..6d13e38aa45 100644 --- a/litellm/llms/anthropic/pass_through/messages/handler.py +++ b/litellm/llms/anthropic/pass_through/messages/handler.py @@ -538,7 +538,7 @@ def anthropic_messages_handler( LiteLLM_Proxy_MCP_Handler, ) - if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools): + if LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(tools=tools): return anthropic_messages_with_mcp( max_tokens=max_tokens, messages=messages, diff --git a/litellm/llms/anthropic/pass_through/messages/mcp_handler.py b/litellm/llms/anthropic/pass_through/messages/mcp_handler.py index 7961a38dedf..41aad100647 100644 --- a/litellm/llms/anthropic/pass_through/messages/mcp_handler.py +++ b/litellm/llms/anthropic/pass_through/messages/mcp_handler.py @@ -83,7 +83,7 @@ async def anthropic_messages_with_mcp( LiteLLM_Proxy_MCP_Handler, ) - mcp_references, other_tools = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools) + mcp_references, other_tools = await LiteLLM_Proxy_MCP_Handler.split_mcp_tools(tools) if not mcp_references: return await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn( @@ -100,7 +100,7 @@ async def anthropic_messages_with_mcp( ( deduplicated_mcp_tools, tool_server_map, - ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + ) = await LiteLLM_Proxy_MCP_Handler.process_mcp_tools_without_openai_transform( context.user_api_key_auth, mcp_references, litellm_trace_id=context.litellm_trace_id, @@ -114,7 +114,7 @@ async def anthropic_messages_with_mcp( ) all_tools: Final = [*anthropic_tools, *(other_tools or ())] - should_auto_execute: Final = LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools( + should_auto_execute: Final = LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools( mcp_tools_with_litellm_proxy=mcp_references ) stream: Final = bool(kwargs.pop("stream", False)) @@ -145,7 +145,7 @@ async def anthropic_messages_with_mcp( if not tool_use_blocks: break - tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_results = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map=tool_server_map, served_tools=deduplicated_mcp_tools, tool_calls=list(tool_use_blocks), diff --git a/litellm/llms/anthropic/pass_through/messages/response_cache.py b/litellm/llms/anthropic/pass_through/messages/response_cache.py index 40dcbd9c8c8..0d8b98962cf 100644 --- a/litellm/llms/anthropic/pass_through/messages/response_cache.py +++ b/litellm/llms/anthropic/pass_through/messages/response_cache.py @@ -94,7 +94,7 @@ class AnthropicMessagesStreamCacheWriter: return self.persisted = True - if not self.caching_handler._should_store_result_in_cache( + if not self.caching_handler.should_store_result_in_cache( original_function=self.caching_handler.original_function, kwargs=self.caching_handler.request_kwargs, ): diff --git a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py index 417017cfb6e..13e1a114bdd 100644 --- a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py @@ -699,7 +699,7 @@ class BaseAnthropicMessagesStreamingIterator: async def _fire_detached_failure_hook(self, exc: Exception) -> None: from litellm._logging import verbose_proxy_logger - on_detached_failure: Final = getattr(self.litellm_logging_obj, "_on_detached_stream_failure", None) + on_detached_failure: Final = getattr(self.litellm_logging_obj, "on_detached_stream_failure", None) if on_detached_failure is None: return try: diff --git a/litellm/llms/anthropic/pass_through/messages/transformation.py b/litellm/llms/anthropic/pass_through/messages/transformation.py index 11698477a33..b35c35f1411 100644 --- a/litellm/llms/anthropic/pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/pass_through/messages/transformation.py @@ -208,7 +208,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): Subclasses whose upstream rejects the role opt in by calling this from their ``transform_anthropic_messages_request``; the first-party Anthropic path forwards ``messages`` untouched and never calls it.""" - from litellm.utils import _supports_factory + from litellm.utils import supports_factory messages: Final = anthropic_messages_request.get("messages") if not isinstance(messages, list): @@ -220,7 +220,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): hoisted: Final = messages[:leading_count] remaining: Final = ( messages[leading_count:] - if _supports_factory( + if supports_factory( model=model, custom_llm_provider=self.custom_llm_provider, key="supports_mid_conversation_system", diff --git a/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py index 83bb146360e..88503908c3c 100644 --- a/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py @@ -74,7 +74,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: ) from litellm.responses.utils import ResponseAPILoggingUtils - chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_usage) + chat_usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(raw_usage) return LiteLLMAnthropicMessagesAdapter._translate_openai_usage_to_anthropic_usage(chat_usage) # ------------------------------------------------------------------ # diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index 5d1a66873a3..ea6e9bdabad 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -22,7 +22,7 @@ from litellm.secret_managers.get_azure_ad_token_provider import ( ) from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams -from litellm.utils import _add_path_to_api_base +from litellm.utils import add_path_to_api_base azure_ad_cache: Final = DualCache() @@ -814,7 +814,7 @@ class BaseAzureLLM(BaseOpenAILLM): # Add the path to the base URL if route not in api_base: - new_url = _add_path_to_api_base(api_base=api_base, ending_path=route) + new_url = add_path_to_api_base(api_base=api_base, ending_path=route) else: new_url = api_base diff --git a/litellm/llms/azure/image_edit/transformation.py b/litellm/llms/azure/image_edit/transformation.py index e4716289a34..05defa6d0ff 100644 --- a/litellm/llms/azure/image_edit/transformation.py +++ b/litellm/llms/azure/image_edit/transformation.py @@ -7,7 +7,7 @@ from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams -from litellm.utils import _add_path_to_api_base +from litellm.utils import add_path_to_api_base class AzureImageEditConfig(OpenAIImageEditConfig): @@ -122,7 +122,7 @@ class AzureImageEditConfig(OpenAIImageEditConfig): # Add the path to the base URL using the model as deployment name if "/openai/deployments/" not in api_base: - new_url = _add_path_to_api_base( + new_url = add_path_to_api_base( api_base=api_base, ending_path=f"/openai/deployments/{model}/images/edits", ) diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index 779a86629e2..fc5e76c9295 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -9,7 +9,7 @@ from httpx import Response import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _audio_or_image_in_message_content, + audio_or_image_in_message_content, convert_content_list_to_str, filter_value_from_dict, ) @@ -27,7 +27,7 @@ from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ModelResponse, ProviderField -from litellm.utils import _add_path_to_api_base, supports_tool_choice +from litellm.utils import add_path_to_api_base, supports_tool_choice if TYPE_CHECKING: from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer @@ -193,9 +193,9 @@ class AzureAIStudioConfig(OpenAIConfig): # Add the path to the base URL if "services.ai.azure.com" in api_base: - new_url = _add_path_to_api_base(api_base=api_base, ending_path="/models/chat/completions") + new_url = add_path_to_api_base(api_base=api_base, ending_path="/models/chat/completions") else: - new_url = _add_path_to_api_base(api_base=api_base, ending_path="/chat/completions") + new_url = add_path_to_api_base(api_base=api_base, ending_path="/chat/completions") # Use the new query_params dictionary final_url: Final = httpx.URL(new_url).copy_with(params=query_params) @@ -245,7 +245,7 @@ class AzureAIStudioConfig(OpenAIConfig): filter_value_from_dict(message_dict, field) # Do nothing if the message contains an image or audio - if _audio_or_image_in_message_content(message): + if audio_or_image_in_message_content(message): continue texts = convert_content_list_to_str(message=message) diff --git a/litellm/llms/azure_ai/image_edit/transformation.py b/litellm/llms/azure_ai/image_edit/transformation.py index 1c626458df4..cf73982bc66 100644 --- a/litellm/llms/azure_ai/image_edit/transformation.py +++ b/litellm/llms/azure_ai/image_edit/transformation.py @@ -9,7 +9,7 @@ from litellm.llms.azure_ai.common_utils import ( ) from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig from litellm.secret_managers.main import get_secret_str -from litellm.utils import _add_path_to_api_base +from litellm.utils import add_path_to_api_base class AzureFoundryFluxImageEditConfig(OpenAIImageEditConfig): @@ -81,12 +81,12 @@ class AzureFoundryFluxImageEditConfig(OpenAIImageEditConfig): # Add the path to the base URL using the model as deployment name # Azure AI Foundry FLUX models use /images/edits for editing if "/openai/deployments/" in api_base: - new_url = _add_path_to_api_base( + new_url = add_path_to_api_base( api_base=api_base, ending_path="/images/edits", ) else: - new_url = _add_path_to_api_base( + new_url = add_path_to_api_base( api_base=api_base, ending_path=f"/openai/deployments/{model}/images/edits", ) diff --git a/litellm/llms/azure_ai/image_generation/cost_calculator.py b/litellm/llms/azure_ai/image_generation/cost_calculator.py index 629c83b7054..beed3881bdf 100644 --- a/litellm/llms/azure_ai/image_generation/cost_calculator.py +++ b/litellm/llms/azure_ai/image_generation/cost_calculator.py @@ -3,15 +3,15 @@ from typing import Final import litellm from litellm.litellm_core_utils.llm_cost_calc.utils import ( - _get_cost_per_unit, calculate_image_response_cost_from_usage, + get_cost_per_unit, resolve_image_model_info, ) from litellm.types.utils import ImageResponse, ModelInfo def _input_cost_per_pixel(resolved: ModelInfo) -> float: - deployment_price: Final = _get_cost_per_unit(resolved, "input_cost_per_pixel", default_value=None) + deployment_price: Final = get_cost_per_unit(resolved, "input_cost_per_pixel", default_value=None) if deployment_price is not None: return deployment_price model_cost_key: Final = resolved.get("key") diff --git a/litellm/llms/azure_ai/rerank/transformation.py b/litellm/llms/azure_ai/rerank/transformation.py index 64372c53f09..24317f8057e 100644 --- a/litellm/llms/azure_ai/rerank/transformation.py +++ b/litellm/llms/azure_ai/rerank/transformation.py @@ -13,7 +13,7 @@ from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers from litellm.llms.cohere.rerank.transformation import CohereRerankConfig from litellm.secret_managers.main import get_secret_str from litellm.types.utils import RerankResponse -from litellm.utils import _add_path_to_api_base +from litellm.utils import add_path_to_api_base class AzureAIRerankConfig(CohereRerankConfig): @@ -52,13 +52,13 @@ class AzureAIRerankConfig(CohereRerankConfig): or normalized_path.endswith("/v2") or normalized_path.endswith("/providers/cohere/v2") ): - return _add_path_to_api_base( + return add_path_to_api_base( api_base=str(original_url.copy_with(path=normalized_path or "/")), ending_path="/rerank", ) # Backwards compatible default: Azure AI rerank was originally exposed under /v1/rerank - return _add_path_to_api_base(api_base=api_base, ending_path="/v1/rerank") + return add_path_to_api_base(api_base=api_base, ending_path="/v1/rerank") def validate_environment( self, diff --git a/litellm/llms/base_llm/base_model_iterator.py b/litellm/llms/base_llm/base_model_iterator.py index 632ef5a24e3..c6bd4e660c6 100644 --- a/litellm/llms/base_llm/base_model_iterator.py +++ b/litellm/llms/base_llm/base_model_iterator.py @@ -98,7 +98,7 @@ class BaseModelResponseIterator: @staticmethod def _string_to_dict_parser(str_line: str) -> dict | None: stripped_json_chunk: dict | None = None - stripped_chunk: Final = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(str_line) + stripped_chunk: Final = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(str_line) try: if stripped_chunk is not None: stripped_json_chunk = json.loads(stripped_chunk) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 886d6d1aea1..38c47b0c8d7 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -28,8 +28,8 @@ from litellm.litellm_core_utils.core_helpers import ( ) from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _parse_content_for_reasoning, drop_lookaround_regex_patterns, + parse_content_for_reasoning, tool_with_sanitized_parameters, ) from litellm.litellm_core_utils.prompt_templates.factory import ( @@ -2387,7 +2387,7 @@ class AmazonConverseConfig(BaseConfig): ( extracted_reasoning_content_str, _content_str, - ) = _parse_content_for_reasoning(content["text"]) + ) = parse_content_for_reasoning(content["text"]) if _content_str is not None: content_str += _content_str if "toolUse" in content: diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_deepseek_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_deepseek_transformation.py index 5699f94d084..41e539e52fd 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_deepseek_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_deepseek_transformation.py @@ -4,7 +4,7 @@ from httpx import Response from litellm import verbose_logger from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _parse_content_for_reasoning, + parse_content_for_reasoning, ) from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( @@ -63,7 +63,7 @@ class AmazonDeepSeekR1Config(AmazonLlamaConfig): message_content: Final = cast(str | None, cast(Choices, response.choices[0]).message.get("content")) if prompt and prompt.strip().endswith("") and message_content: message_content_with_reasoning_token: Final = "" + message_content - reasoning, content = _parse_content_for_reasoning(message_content_with_reasoning_token) + reasoning, content = parse_content_for_reasoning(message_content_with_reasoning_token) provider_specific_fields: Final = cast(Choices, response.choices[0]).message.provider_specific_fields or {} if reasoning: provider_specific_fields["reasoning_content"] = reasoning diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 238b1ccaa30..0352234ecb8 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -246,9 +246,9 @@ def convert_bedrock_invoke_output_format_to_inline_schema( def _bedrock_model_supports(model: str, key: str) -> bool: - from litellm.utils import _supports_factory + from litellm.utils import supports_factory - return _supports_factory(model=model, custom_llm_provider="bedrock", key=key) + return supports_factory(model=model, custom_llm_provider="bedrock", key=key) def apply_bedrock_invoke_structured_output( diff --git a/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py b/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py index acb0cc8dcb7..463511ff333 100644 --- a/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py +++ b/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py @@ -25,9 +25,9 @@ from litellm.types.images.main import ImageEditOptionalRequestParams from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import FileTypes, ImageObject, ImageResponse from litellm.utils import ( - _get_model_cost_key, - _get_potential_model_names, + get_model_cost_key, get_model_info, + get_potential_model_names, ) if TYPE_CHECKING: @@ -192,7 +192,7 @@ def _supports_nova_canvas_image_edit_from_model_cost(model: str) -> bool: pass try: - potential: Final = _get_potential_model_names(model=model, custom_llm_provider=None) + potential: Final = get_potential_model_names(model=model, custom_llm_provider=None) for field in ( "combined_model_name", "combined_stripped_model_name", @@ -206,7 +206,7 @@ def _supports_nova_canvas_image_edit_from_model_cost(model: str) -> bool: pass for name in candidates: - key = _get_model_cost_key(name) + key = get_model_cost_key(name) if key is None: continue entry = _litellm.model_cost.get(key) or {} diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index eb314450f08..2ceac5899fd 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -16,7 +16,7 @@ from typing import Final, NoReturn, Protocol, runtime_checkable from pydantic import JsonValue, TypeAdapter import litellm -from litellm._logging import _redact_string, verbose_proxy_logger +from litellm._logging import redact_string, verbose_proxy_logger from litellm.constants import ( BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY, BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY, @@ -192,9 +192,9 @@ def _pending_session_update(scope: Mapping[str, object]) -> str | None: def _raise_provider_failure(scope: MutableMapping[str, object], failure: BaseException) -> NoReturn: error: Final = _as_bedrock_error(failure) - verbose_proxy_logger.error("Bedrock Realtime: provider stream failed: %s", _redact_string(str(error))) + verbose_proxy_logger.error("Bedrock Realtime: provider stream failed: %s", redact_string(str(error))) if scope.get(BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY) is True: - scope[BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY] = _redact_string(str(error)) + scope[BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY] = redact_string(str(error)) raise error from failure diff --git a/litellm/llms/chatgpt/responses/transformation.py b/litellm/llms/chatgpt/responses/transformation.py index d66e3d97d66..ae998238cd0 100644 --- a/litellm/llms/chatgpt/responses/transformation.py +++ b/litellm/llms/chatgpt/responses/transformation.py @@ -6,7 +6,7 @@ import httpx from litellm.exceptions import AuthenticationError from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _safe_convert_created_field, + safe_convert_created_field, ) from litellm.llms.openai.common_utils import OpenAIError from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig @@ -216,7 +216,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): if not response_payload.get("output") and streamed_output_items: response_payload["output"] = [item for _, item in sorted(streamed_output_items.items())] if "created_at" in response_payload: - response_payload["created_at"] = _safe_convert_created_field(response_payload["created_at"]) + response_payload["created_at"] = safe_convert_created_field(response_payload["created_at"]) try: return ResponsesAPIResponse(**response_payload) except Exception: diff --git a/litellm/llms/codestral/completion/transformation.py b/litellm/llms/codestral/completion/transformation.py index baa134bb398..7d675b12f4e 100644 --- a/litellm/llms/codestral/completion/transformation.py +++ b/litellm/llms/codestral/completion/transformation.py @@ -83,7 +83,7 @@ class CodestralTextCompletionConfig(OpenAITextCompletionConfig): finish_reason = None logprobs = None - chunk_data = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk_data) or "" + chunk_data = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(chunk_data) or "" chunk_data = chunk_data.strip() if len(chunk_data) == 0 or chunk_data == "[DONE]": return { diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 53ea85a2740..202bfef0d8a 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -33,7 +33,7 @@ import litellm import litellm.litellm_core_utils import litellm.types import litellm.types.utils -from litellm._logging import _redact_string, verbose_logger +from litellm._logging import redact_string, verbose_logger from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta from litellm.constants import MAX_FILE_LIST_LIMIT, REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES from litellm.files.types import FileContentStreamingResult @@ -411,12 +411,12 @@ def _mask_presigned_request_headers(transformed_request: bytes | str | dict) -> return transformed_request from litellm.litellm_core_utils.litellm_logging import ( - _get_masked_values, # pyright: ignore[reportPrivateUsage] # the shared header-masking helper has no public name + get_masked_values, ) return { **transformed_request, - "headers": _get_masked_values(request_headers), + "headers": get_masked_values(request_headers), } @@ -6342,7 +6342,7 @@ class BaseLLMHTTPHandler: if provider_config.requires_session_configuration(): _session_config = provider_config.session_configuration_request(model) if _session_config: - _session_config = realtime_streaming._maybe_inject_guardrail_auto_response_disable( + _session_config = realtime_streaming.maybe_inject_guardrail_auto_response_disable( _session_config ) await backend_ws.send(_session_config) @@ -6364,7 +6364,7 @@ class BaseLLMHTTPHandler: # success_handler / async_success_handler payloads. realtime_streaming.store_message(synthetic_session_str) await websocket.send_text(synthetic_session_str) - realtime_streaming._session_created_sent_to_client = True + realtime_streaming.session_created_sent_to_client = True verbose_logger.debug("Sent synthetic session.created to client to unblock connection") await realtime_streaming.bidirectional_forward() @@ -6374,7 +6374,7 @@ class BaseLLMHTTPHandler: await close_after_upstream_handshake_refusal(websocket, e.response.status_code) except Exception as e: verbose_logger.exception("Error connecting to backend: %s", e) - redacted_error: Final = _redact_string(str(e)) + redacted_error: Final = redact_string(str(e)) try: await websocket.send_text(realtime_error_event(redacted_error, error_type="server_error")) except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below @@ -6383,7 +6383,7 @@ class BaseLLMHTTPHandler: await websocket.close( code=1011, reason=websocket_close_reason( - _redact_string(f"Internal server error: {e}"), + redact_string(f"Internal server error: {e}"), fallback="Internal server error", ), ) @@ -6790,7 +6790,7 @@ class BaseLLMHTTPHandler: except Exception as e: verbose_logger.exception("Error in responses WS: %s", e) try: - await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {e}")) + await websocket.close(code=1011, reason=redact_string(f"Internal server error: {e}")) except RuntimeError as close_error: if "already completed" in str(close_error) or "websocket.close" in str(close_error): pass diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index d669f2acc6d..5fd553126cc 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -11,11 +11,11 @@ from pydantic import BaseModel from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _handle_invalid_parallel_tool_calls, - _should_convert_tool_call_to_json_mode, + handle_invalid_parallel_tool_calls, + should_convert_tool_call_to_json_mode, ) from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _extract_reasoning_content, # pyright: ignore[reportPrivateUsage] # same import as the OpenAI transformation + extract_reasoning_content, merge_consecutive_system_messages, strip_litellm_internal_message_fields, strip_name_from_message, @@ -561,7 +561,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): content_str: Final = DatabricksConfig.extract_content_str(message["content"]) if block_reasoning_content is not None: return block_reasoning_content, content_str - return _extract_reasoning_content({**message, "content": content_str}) + return extract_reasoning_content({**message, "content": content_str}) @staticmethod def extract_citations( @@ -588,14 +588,14 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): for _tc in tool_calls: _openai_tc = ChatCompletionMessageToolCall(**_tc) _openai_tool_calls.append(_openai_tc) - fixed_tool_calls = _handle_invalid_parallel_tool_calls(_openai_tool_calls) + fixed_tool_calls = handle_invalid_parallel_tool_calls(_openai_tool_calls) if fixed_tool_calls is not None: tool_calls = fixed_tool_calls translated_message: Message | None = None finish_reason: str | None = None - if tool_calls and _should_convert_tool_call_to_json_mode( + if tool_calls and should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=json_mode, ): diff --git a/litellm/llms/databricks/streaming_utils.py b/litellm/llms/databricks/streaming_utils.py index 92f82a3f8d7..c729e6fcc00 100644 --- a/litellm/llms/databricks/streaming_utils.py +++ b/litellm/llms/databricks/streaming_utils.py @@ -105,7 +105,7 @@ class ModelResponseIterator: raise RuntimeError(f"Error receiving chunk from stream: {e}") try: - chunk = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or "" + chunk = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(chunk) or "" chunk = chunk.strip() if len(chunk) > 0: json_chunk: Final = json.loads(chunk) @@ -150,7 +150,7 @@ class ModelResponseIterator: raise RuntimeError(f"Error receiving chunk from stream: {e}") try: - chunk = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or "" + chunk = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(chunk) or "" chunk = chunk.strip() if chunk == "[DONE]": raise StopAsyncIteration diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index e64cbf88d95..3b926ccd1c2 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -1070,7 +1070,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): else: _chat_completion_usage = get_empty_usage() - responses_api_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + responses_api_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( _chat_completion_usage, ) _usage_dict: Final = responses_api_usage.model_dump() @@ -1496,7 +1496,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): else: _tool_call_chat_completion_usage = get_empty_usage() tool_call_responses_api_usage = ( - LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( _tool_call_chat_completion_usage, ) ) diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index 8b85b668cba..7bdf253f28f 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -23,7 +23,7 @@ from litellm.types.llms.openai import ( ) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import LlmProviders -from litellm.utils import _cached_get_model_info_helper +from litellm.utils import cached_get_model_info_helper from ..authenticator import Authenticator from ..common_utils import ( @@ -53,7 +53,7 @@ def github_copilot_supports_responses_api(model: str) -> bool: register_model, which also clears the cache used here). """ try: - info: Final = _cached_get_model_info_helper(model=model, custom_llm_provider="github_copilot") + info: Final = cached_get_model_info_helper(model=model, custom_llm_provider="github_copilot") except Exception as e: verbose_logger.debug( "github_copilot_supports_responses_api: get_model_info failed for %s: %s", diff --git a/litellm/llms/hosted_vllm/chat/transformation.py b/litellm/llms/hosted_vllm/chat/transformation.py index 43eb2af171e..58ab297a198 100644 --- a/litellm/llms/hosted_vllm/chat/transformation.py +++ b/litellm/llms/hosted_vllm/chat/transformation.py @@ -7,9 +7,9 @@ from collections.abc import Coroutine from typing import Final, Literal, cast, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _get_image_mime_type_from_url, + get_image_mime_type_from_url, ) -from litellm.litellm_core_utils.prompt_templates.factory import _parse_mime_type +from litellm.litellm_core_utils.prompt_templates.factory import parse_mime_type from litellm.litellm_core_utils.reasoning_effort_utils import ( reasoning_effort_from_thinking_budget, ) @@ -23,7 +23,7 @@ from litellm.types.llms.openai import ( ChatCompletionVideoUrlObject, ) -from ....utils import _remove_additional_properties, _remove_strict_from_schema +from ....utils import remove_additional_properties, remove_strict_from_schema from ...openai.chat.gpt_transformation import OpenAIGPTConfig @@ -88,8 +88,8 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): ) -> dict: _tools = non_default_params.pop("tools", None) if _tools is not None: - _tools = _remove_additional_properties(_tools) - _tools = _remove_strict_from_schema(_tools) + _tools = remove_additional_properties(_tools) + _tools = remove_strict_from_schema(_tools) if isinstance(_tools, list): _tools = self._convert_custom_tools_to_function_tools(_tools) if _tools is not None: @@ -122,11 +122,11 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): if format and format.startswith("video/"): return True elif file_data: - mime_type = _parse_mime_type(file_data) + mime_type = parse_mime_type(file_data) if mime_type and mime_type.startswith("video/"): return True elif file_id: - mime_type = _get_image_mime_type_from_url(file_id) + mime_type = get_image_mime_type_from_url(file_id) if mime_type and mime_type.startswith("video/"): return True return False diff --git a/litellm/llms/manus/responses/transformation.py b/litellm/llms/manus/responses/transformation.py index f5733a152ba..775ae3e3ad8 100644 --- a/litellm/llms/manus/responses/transformation.py +++ b/litellm/llms/manus/responses/transformation.py @@ -7,7 +7,7 @@ import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _safe_convert_created_field, + safe_convert_created_field, ) from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.openai.common_utils import OpenAIError @@ -181,11 +181,11 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): # Manus uses camelCase "createdAt" instead of snake_case "created_at" if "createdAt" in raw_response_json and "created_at" not in raw_response_json: - raw_response_json["created_at"] = _safe_convert_created_field(raw_response_json["createdAt"]) + raw_response_json["created_at"] = safe_convert_created_field(raw_response_json["createdAt"]) # Ensure created_at is set if "created_at" in raw_response_json: - raw_response_json["created_at"] = _safe_convert_created_field(raw_response_json["created_at"]) + raw_response_json["created_at"] = safe_convert_created_field(raw_response_json["created_at"]) except Exception: raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code) @@ -273,11 +273,11 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig): # Manus uses camelCase "createdAt" instead of snake_case "created_at" if "createdAt" in raw_response_json and "created_at" not in raw_response_json: - raw_response_json["created_at"] = _safe_convert_created_field(raw_response_json["createdAt"]) + raw_response_json["created_at"] = safe_convert_created_field(raw_response_json["createdAt"]) # Ensure created_at is set if "created_at" in raw_response_json: - raw_response_json["created_at"] = _safe_convert_created_field(raw_response_json["created_at"]) + raw_response_json["created_at"] = safe_convert_created_field(raw_response_json["created_at"]) except Exception: raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code) diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index b78797c35a3..b9d70244b0a 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -394,12 +394,12 @@ class MistralConfig(OpenAIGPTConfig): import copy from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH - from litellm.utils import _remove_json_schema_refs + from litellm.utils import remove_json_schema_refs cleaned_tools = copy.deepcopy(tools) # Apply all cleaning functions with max_depth protection - cleaned_tools = _remove_json_schema_refs(cleaned_tools, max_depth=DEFAULT_MAX_RECURSE_DEPTH) + cleaned_tools = remove_json_schema_refs(cleaned_tools, max_depth=DEFAULT_MAX_RECURSE_DEPTH) return cleaned_tools diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index dc3705b0fe5..132d3a38536 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -10,9 +10,9 @@ import litellm from litellm._uuid import uuid from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _extract_reasoning_content, convert_content_list_to_str, extract_images_from_message, + extract_reasoning_content, ) from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException @@ -266,7 +266,7 @@ class OllamaChatConfig(BaseConfig): ) ) new_tools.append(ollama_tool_call) - reasoning_content, parsed_content = _extract_reasoning_content(cast(dict, m)) + reasoning_content, parsed_content = extract_reasoning_content(cast(dict, m)) content_str = convert_content_list_to_str(cast(AllMessageValues, m)) images = extract_images_from_message(cast(AllMessageValues, m)) @@ -349,10 +349,10 @@ class OllamaChatConfig(BaseConfig): elif response_json_message.get("content") is not None: # parse reasoning content from content from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _parse_content_for_reasoning, + parse_content_for_reasoning, ) - reasoning_content, content = _parse_content_for_reasoning(response_json_message["content"]) + reasoning_content, content = parse_content_for_reasoning(response_json_message["content"]) response_json_message["reasoning_content"] = reasoning_content response_json_message["content"] = content diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index 9687378e5c3..2a17456ffa8 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -66,14 +66,14 @@ class _OllamaGenerateReasoning(LiteLLMBaseModel): """Reasoning reaches `/api/generate` either in the top-level `thinking` field or inline in `` tags, never both. The field wins, matching `ollama_chat`.""" from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _parse_content_for_reasoning, + parse_content_for_reasoning, ) if self.thinking: return self.thinking, self.response if self.response is None: return None, None - return _parse_content_for_reasoning(self.response) + return parse_content_for_reasoning(self.response) class OllamaConfig(BaseConfig): diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index bf6b52225f2..97e412942aa 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -5,9 +5,9 @@ from typing import Final import litellm from litellm.utils import ( - _supports_factory, declared_value_factory, is_explicitly_disabled_factory, + supports_factory, ) from .gpt_transformation import OpenAIGPTConfig @@ -157,7 +157,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig): the shared ``_supports_factory`` helper. Returns False for unknown models (safe fallback). """ - return _supports_factory( + return supports_factory( model=cls._model_map_lookup_name(model), custom_llm_provider=None, key=f"supports_{level}_reasoning_effort", diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 2dfb397f427..6380c9390c6 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -15,9 +15,9 @@ import litellm from litellm.constants import OPENAI_SYSTEM_MESSAGES_FIRST_PROVIDERS from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _extract_reasoning_content, - _handle_invalid_parallel_tool_calls, - _should_convert_tool_call_to_json_mode, + extract_reasoning_content, + handle_invalid_parallel_tool_calls, + should_convert_tool_call_to_json_mode, ) from litellm.litellm_core_utils.prompt_templates.common_utils import ( drop_non_python_regex_patterns, @@ -609,7 +609,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): for _tc in tool_calls: _openai_tc = chat_completion_tool_call_from_dict(_tc) _openai_tool_calls.append(_openai_tc) - fixed_tool_calls = _handle_invalid_parallel_tool_calls(_openai_tool_calls) + fixed_tool_calls = handle_invalid_parallel_tool_calls(_openai_tool_calls) if fixed_tool_calls is not None: new_tool_calls = fixed_tool_calls @@ -621,7 +621,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): translated_message: Message | None = None finish_reason: str | None = None - if new_tool_calls and _should_convert_tool_call_to_json_mode( + if new_tool_calls and should_convert_tool_call_to_json_mode( tool_calls=new_tool_calls, convert_tool_call_to_json_mode=json_mode, ): @@ -636,7 +636,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): ( reasoning_content, content_str, - ) = _extract_reasoning_content(cast(dict, choice["message"])) + ) = extract_reasoning_content(cast(dict, choice["message"])) translated_message = Message( role="assistant", diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index b4a94a1c0bf..182b24c9901 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -7,7 +7,7 @@ This requires websockets, and is currently only supported on LiteLLM Proxy. import ssl from typing import Any, Final, cast -from litellm._logging import _redact_string, verbose_logger +from litellm._logging import redact_string, verbose_logger from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES from litellm.types.realtime import RealtimeQueryParams @@ -190,7 +190,7 @@ class OpenAIRealtime(OpenAIChatCompletion): await close_after_upstream_handshake_refusal(websocket, e.response.status_code) except Exception as e: try: - await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {e}")) + await websocket.close(code=1011, reason=redact_string(f"Internal server error: {e}")) except RuntimeError as close_error: if "already completed" in str(close_error) or "websocket.close" in str(close_error): # The WebSocket is already closed or the response is completed, so we can ignore this error diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 53f57456126..e7463fb66f6 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -13,7 +13,7 @@ from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _safe_convert_created_field, + safe_convert_created_field, ) from litellm.litellm_core_utils.prompt_templates.common_utils import ( drop_non_python_regex_patterns, @@ -157,10 +157,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): @staticmethod def _supports_reasoning_param(model: str) -> bool: - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper try: - info: Final = _get_model_info_helper( + info: Final = get_model_info_helper( model=model.split("/")[-1], custom_llm_provider=LlmProviders.OPENAI.value ) except Exception: @@ -579,7 +579,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): additional_args={"complete_input_dict": {}}, ) raw_response_json: Final = _RAW_RESPONSE_JSON.validate_python(raw_response.json()) - raw_response_json["created_at"] = _safe_convert_created_field(raw_response_json["created_at"]) + raw_response_json["created_at"] = safe_convert_created_field(raw_response_json["created_at"]) except Exception: raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code) raw_response_headers: Final = dict(raw_response.headers) @@ -979,7 +979,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): additional_args={"complete_input_dict": {}}, ) raw_response_json: Final = _RAW_RESPONSE_JSON.validate_python(raw_response.json()) - raw_response_json["created_at"] = _safe_convert_created_field(raw_response_json["created_at"]) + raw_response_json["created_at"] = safe_convert_created_field(raw_response_json["created_at"]) except Exception: raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code) raw_response_headers: Final = dict(raw_response.headers) diff --git a/litellm/llms/openai_like/model_info.py b/litellm/llms/openai_like/model_info.py index 24135287ae9..e3b2dc427c4 100644 --- a/litellm/llms/openai_like/model_info.py +++ b/litellm/llms/openai_like/model_info.py @@ -11,7 +11,7 @@ from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.llms.base import LiteLLMBaseModel -from litellm.utils import _add_path_to_api_base # pyright: ignore[reportPrivateUsage] # shared provider URL helper +from litellm.utils import add_path_to_api_base MODEL_INFO_REFRESH_SECONDS: Final = 300 MODEL_INFO_REFRESH_CONCURRENCY: Final = 8 @@ -66,7 +66,7 @@ async def get_openai_compatible_model_info( client: AsyncHTTPHandler, cache: InMemoryCache, ) -> Mapping[str, int]: - url: Final = _add_path_to_api_base(api_base, "/v1/models") + url: Final = add_path_to_api_base(api_base, "/v1/models") cache_key: Final = ( "upstream_model_info:" + hashlib.sha256(json.dumps((url, sorted(headers.items()))).encode()).hexdigest() ) diff --git a/litellm/llms/sagemaker/common_utils.py b/litellm/llms/sagemaker/common_utils.py index 0d86752fc5f..1f7fba3c4c4 100644 --- a/litellm/llms/sagemaker/common_utils.py +++ b/litellm/llms/sagemaker/common_utils.py @@ -105,7 +105,7 @@ class AWSEventStreamDecoder: message = self._parse_message_from_event(event) if message: # remove data: prefix and "\n\n" at the end - message = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(message) or "" + message = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(message) or "" message = message.replace("\n\n", "") # Accumulate JSON data @@ -154,7 +154,7 @@ class AWSEventStreamDecoder: if message: verbose_logger.debug("sagemaker parsed chunk bytes %s", message) # remove data: prefix and "\n\n" at the end - message = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(message) or "" + message = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(message) or "" message = message.replace("\n\n", "") # Accumulate JSON data diff --git a/litellm/llms/vertex_ai/cost_calculator.py b/litellm/llms/vertex_ai/cost_calculator.py index 8b00fc2e925..864619f91b0 100644 --- a/litellm/llms/vertex_ai/cost_calculator.py +++ b/litellm/llms/vertex_ai/cost_calculator.py @@ -5,9 +5,9 @@ from typing import Final, Literal import litellm from litellm import verbose_logger from litellm.litellm_core_utils.llm_cost_calc.utils import ( - _is_above_128k, generic_cost_per_token, get_vertex_regional_endpoint_uplift, + is_above_128k, ) from litellm.types.utils import ModelInfo, Usage @@ -100,7 +100,7 @@ def cost_per_character( else: try: if ( - _is_above_128k(tokens=prompt_characters * 4) # 1 token = 4 char + is_above_128k(tokens=prompt_characters * 4) # 1 token = 4 char and model not in models_without_dynamic_pricing ): ## check if character pricing, else default to token pricing @@ -142,7 +142,7 @@ def cost_per_character( completion_tokens: Final = usage.completion_tokens try: if ( - _is_above_128k(tokens=completion_characters * 4) # 1 token = 4 char + is_above_128k(tokens=completion_characters * 4) # 1 token = 4 char and model not in models_without_dynamic_pricing ): assert ( @@ -186,14 +186,14 @@ def _handle_128k_pricing( prompt_tokens: Final = usage.prompt_tokens completion_tokens: Final = usage.completion_tokens - if _is_above_128k(tokens=prompt_tokens) and input_cost_per_token_above_128k_tokens is not None: + if is_above_128k(tokens=prompt_tokens) and input_cost_per_token_above_128k_tokens is not None: prompt_cost = prompt_tokens * input_cost_per_token_above_128k_tokens else: prompt_cost = prompt_tokens * (model_info["input_cost_per_token"] or 0.0) ## CALCULATE OUTPUT COST output_cost_per_token_above_128k_tokens = model_info.get("output_cost_per_token_above_128k_tokens") - if _is_above_128k(tokens=completion_tokens) and output_cost_per_token_above_128k_tokens is not None: + if is_above_128k(tokens=completion_tokens) and output_cost_per_token_above_128k_tokens is not None: completion_cost = completion_tokens * output_cost_per_token_above_128k_tokens else: completion_cost = completion_tokens * (model_info["output_cost_per_token"] or 0.0) diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index d641d3da454..51b87233c30 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -17,14 +17,14 @@ import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _get_image_mime_type_from_url, + get_image_mime_type_from_url, ) from litellm.litellm_core_utils.prompt_templates.factory import ( - _get_thought_signature_from_tool, convert_generic_image_chunk_to_openai_image_obj, convert_to_anthropic_image_obj, convert_to_gemini_tool_call_invoke, convert_to_gemini_tool_call_result, + get_thought_signature_from_tool, response_schema_prompt, ) from litellm.litellm_core_utils.prompt_templates.image_handling import RemoteMedia, async_inline_remote_media @@ -570,7 +570,7 @@ def _process_gemini_media( file_data = cast(FileDataType, {"file_uri": image_url}) part = {"file_data": file_data} return _apply_gemini_metadata(part, model, media_resolution_enum, video_metadata) - elif "https://" in image_url and (image_type := format or _get_image_mime_type_from_url(image_url)) is not None: + elif "https://" in image_url and (image_type := format or get_image_mime_type_from_url(image_url)) is not None: file_data = FileDataType(mime_type=image_type, file_uri=image_url) part = {"file_data": file_data} return _apply_gemini_metadata(part, model, media_resolution_enum, video_metadata) @@ -658,13 +658,13 @@ def _collect_tool_call_thought_signatures( for tool in tool_calls: if not isinstance(tool, dict): continue - signature = _get_thought_signature_from_tool(tool) + signature = get_thought_signature_from_tool(tool) if signature: signatures += (signature,) function_call: Final = assistant_msg.get("function_call") if isinstance(function_call, dict): - signature = _get_thought_signature_from_tool({"function": function_call}) + signature = get_thought_signature_from_tool({"function": function_call}) if signature: signatures += (signature,) @@ -1301,7 +1301,7 @@ def _vertex_inlines(media: RemoteMedia) -> bool: if media.url.startswith(GEMINI_FILES_API_URI_PREFIX): return False return media.url.startswith("http://") or ( - _explicit_mime_type(media.fields) is None and _get_image_mime_type_from_url(media.url) is None + _explicit_mime_type(media.fields) is None and get_image_mime_type_from_url(media.url) is None ) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index b5f32d57061..b8a70c63a4c 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -10,6 +10,7 @@ from functools import partial from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args import httpx +from pydantic import JsonValue import litellm from litellm import verbose_logger @@ -27,7 +28,7 @@ from litellm.constants import ( from litellm.exceptions import UnsupportedParamsError from litellm.litellm_core_utils.json_fragment_accumulator import JSONFragmentAccumulator from litellm.litellm_core_utils.prompt_templates.factory import ( - _encode_tool_call_id_with_signature, + encode_tool_call_id_with_signature, ) from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.custom_httpx.http_handler import ( @@ -85,7 +86,7 @@ from litellm.utils import ( supports_reasoning, ) -from ....utils import _remove_additional_properties, _remove_strict_from_schema +from ....utils import remove_additional_properties, remove_strict_from_schema from ..common_utils import ( VertexAIError, _build_json_schema, @@ -611,9 +612,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): google_maps_retrieval_config: dict | None = None computerUse: dict | None = None # remove 'additionalProperties' from tools - value = _remove_additional_properties(value) + schema_value: Final[JsonValue] = cast(JsonValue, value) # cast-ok: tool schemas come from JSON request data + remove_additional_properties(schema_value) # remove 'strict' from tools - value = _remove_strict_from_schema(value) + remove_strict_from_schema(schema_value) for tool in value: openai_function_object: ChatCompletionToolParamFunctionChunk | None = None @@ -778,7 +780,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def apply_response_schema_transformation(self, value: dict, optional_params: dict, model: str): new_value = deepcopy(value) # remove 'strict' from json schema (not supported by Gemini) - new_value = _remove_strict_from_schema(new_value) + new_value = remove_strict_from_schema(new_value) # Automatically use responseJsonSchema for Gemini 2.0+ models # responseJsonSchema uses standard JSON Schema format and supports additionalProperties @@ -787,7 +789,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if not use_json_schema: # For responseSchema, remove 'additionalProperties' (not supported) - new_value = _remove_additional_properties(new_value) + new_value = remove_additional_properties(new_value) # Handle response type if new_value.get("type") == "json_object": @@ -1591,7 +1593,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Embed thought signature in ID for OpenAI client compatibility if thought_signature: _tool_response_chunk["provider_specific_fields"] = {"thought_signature": thought_signature} - _tool_response_chunk["id"] = _encode_tool_call_id_with_signature( + _tool_response_chunk["id"] = encode_tool_call_id_with_signature( _tool_response_chunk["id"] or "", thought_signature ) _tools.append(_tool_response_chunk) @@ -3307,7 +3309,7 @@ class ModelResponseIterator: return self.chunk_parser(chunk=json_chunk) def handle_accumulated_json_chunk(self, chunk: str, is_final: bool = False) -> Optional["ModelResponseStream"]: - message: Final = (litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or "").replace("\n\n", "") + message: Final = (litellm.CustomStreamWrapper.strip_sse_data_from_chunk(chunk) or "").replace("\n\n", "") self._json_buffer.append(message) # Mid-stream, defer parsing until the buffer's last byte can close a value: @@ -3336,7 +3338,7 @@ class ModelResponseIterator: def _common_chunk_parsing_logic(self, chunk: str) -> Optional["ModelResponseStream"]: try: - chunk = litellm.CustomStreamWrapper._strip_sse_data_from_chunk(chunk) or "" + chunk = litellm.CustomStreamWrapper.strip_sse_data_from_chunk(chunk) or "" if len(chunk) > 0: """ Check if initial chunk valid json diff --git a/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py b/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py index 2721f73207c..da77fa338db 100644 --- a/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py +++ b/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py @@ -19,7 +19,7 @@ from litellm.types.utils import ( PromptTokensDetailsWrapper, Usage, ) -from litellm.utils import _count_characters, is_base64_encoded +from litellm.utils import count_characters, is_base64_encoded from ...base_llm.embedding.transformation import BaseEmbeddingConfig from ..common_utils import VertexAIError @@ -245,7 +245,7 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): prompt += text if prompt is not None: - character_count = _count_characters(prompt) + character_count = count_characters(prompt) ## Calculate image embeddings usage image_count = 0 diff --git a/litellm/llms/vllm/common_utils.py b/litellm/llms/vllm/common_utils.py index 2da2269e1f9..6484cdc88e3 100644 --- a/litellm/llms/vllm/common_utils.py +++ b/litellm/llms/vllm/common_utils.py @@ -7,7 +7,7 @@ from litellm.llms.base_llm.base_utils import BaseLLMModelInfo from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues -from litellm.utils import _add_path_to_api_base +from litellm.utils import add_path_to_api_base class VLLMError(BaseLLMException): @@ -69,7 +69,7 @@ class VLLMModelInfo(BaseLLMModelInfo): "VLLM_API_BASE or VLLM_API_KEY is not set. Please set the environment variable, to query VLLM's `/models` endpoint." ) - url: Final = _add_path_to_api_base(api_base, endpoint) + url: Final = add_path_to_api_base(api_base, endpoint) response: Final = litellm.module_level_client.get( url=url, ) diff --git a/litellm/llms/volcengine/responses/transformation.py b/litellm/llms/volcengine/responses/transformation.py index 82c4f7f61a7..89a54e67c0d 100644 --- a/litellm/llms/volcengine/responses/transformation.py +++ b/litellm/llms/volcengine/responses/transformation.py @@ -9,7 +9,7 @@ import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _safe_convert_created_field, + safe_convert_created_field, ) from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig @@ -245,7 +245,7 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig): ) raw_response_json: Final = self._parsed_response_body(raw_response) if "created_at" in raw_response_json: - raw_response_json["created_at"] = _safe_convert_created_field(raw_response_json["created_at"]) + raw_response_json["created_at"] = safe_convert_created_field(raw_response_json["created_at"]) except Exception: raise VolcEngineError(message=raw_response.text, status_code=raw_response.status_code) diff --git a/litellm/llms/watsonx/chat/transformation.py b/litellm/llms/watsonx/chat/transformation.py index cc616ab5f9f..4ff2212711d 100644 --- a/litellm/llms/watsonx/chat/transformation.py +++ b/litellm/llms/watsonx/chat/transformation.py @@ -13,7 +13,7 @@ from litellm.types.llms.watsonx import ( WatsonXModelPattern, ) -from ....utils import _remove_additional_properties, _remove_strict_from_schema +from ....utils import remove_additional_properties, remove_strict_from_schema from ...openai.chat.gpt_transformation import OpenAIGPTConfig from ..common_utils import IBMWatsonXMixin @@ -56,9 +56,9 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig): _tools = non_default_params.pop("tools", None) if _tools is not None: # remove 'additionalProperties' from tools - _tools = _remove_additional_properties(_tools) + _tools = remove_additional_properties(_tools) # remove 'strict' from tools - _tools = _remove_strict_from_schema(_tools) + _tools = remove_strict_from_schema(_tools) if _tools is not None: non_default_params["tools"] = _tools diff --git a/litellm/main.py b/litellm/main.py index 5b0af2e479c..4f95ce33191 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -28,7 +28,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, Union, cast, get_args from urllib.parse import urlsplit -from litellm._logging import _redact_string +from litellm._logging import redact_string from litellm._uuid import uuid if TYPE_CHECKING: @@ -93,8 +93,8 @@ from litellm.litellm_core_utils.get_provider_specific_headers import ( ProviderSpecificHeaderUtils, ) from litellm.litellm_core_utils.health_check_utils import ( - _create_health_check_response, _filter_model_params, + create_health_check_response, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.mock_functions import ( @@ -131,10 +131,10 @@ from litellm.llms.vertex_ai.common_utils import ( VertexAIModelRoute, get_vertex_ai_model_route, ) -from litellm.realtime_api.main import _realtime_health_check +from litellm.realtime_api.main import realtime_health_check from litellm.secret_managers.main import get_secret_bool, get_secret_str from litellm.types.completion import ( - _CompletionDispatchContext, + CompletionDispatchContext, _CompletionDispatchResult, ) from litellm.types.litellm_params import ControlOptions, RetryStrategy @@ -158,7 +158,6 @@ from litellm.utils import ( TextCompletionStreamWrapper, TranscriptionResponse, Usage, - _get_model_info_helper, add_provider_specific_params_to_optional_params, async_mock_completion_streaming_obj, convert_to_model_response_object, @@ -166,6 +165,7 @@ from litellm.utils import ( create_tokenizer, get_llm_provider, get_model_info, + get_model_info_helper, get_non_default_completion_params, get_non_default_transcription_params, get_optional_params_embeddings, @@ -1104,7 +1104,7 @@ def responses_api_bridge_check( try: model_info = cast( dict, - _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider), + get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider), ) if model_info.get("mode") is None and model.startswith("responses/"): model = model.replace("responses/", "") @@ -1320,25 +1320,25 @@ def _register_custom_pricing_for_request( ) -def _dispatch_metadata(ctx: _CompletionDispatchContext) -> Mapping[str, object] | None: +def _dispatch_metadata(ctx: CompletionDispatchContext) -> Mapping[str, object] | None: return ctx.metadata -def _dispatch_client_http(ctx: _CompletionDispatchContext) -> HTTPHandler | AsyncHTTPHandler | None: +def _dispatch_client_http(ctx: CompletionDispatchContext) -> HTTPHandler | AsyncHTTPHandler | None: return ctx.client def _dispatch_client_azure( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> openai.AzureOpenAI | openai.AsyncAzureOpenAI | HTTPHandler | AsyncHTTPHandler | None: return ctx.client -def _dispatch_client_openai(ctx: _CompletionDispatchContext) -> openai.OpenAI | openai.AsyncOpenAI | None: +def _dispatch_client_openai(ctx: CompletionDispatchContext) -> openai.OpenAI | openai.AsyncOpenAI | None: return ctx.client -def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_azure(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: _azure_detection_model: Final = ctx._azure_detection_model acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -1472,7 +1472,7 @@ def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResul return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract -def _complete_azure_text(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_azure_text(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -1566,7 +1566,7 @@ def _complete_azure_text(ctx: _CompletionDispatchContext) -> _CompletionDispatch return response -def _complete_deepseek(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_deepseek(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -1617,7 +1617,7 @@ def _complete_deepseek(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe return response -def _complete_azure_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_azure_ai(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -1772,7 +1772,7 @@ def _complete_azure_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe def _complete_text_completion_openai( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -1854,7 +1854,7 @@ def _complete_text_completion_openai( def _complete_fireworks_ai( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base @@ -1910,7 +1910,7 @@ def _complete_fireworks_ai( return response -def _complete_together_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_together_ai(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -1960,7 +1960,7 @@ def _complete_together_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatc return response -def _complete_heroku(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_heroku(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -2010,7 +2010,7 @@ def _complete_heroku(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu return response -def _complete_ragflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_ragflow(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -2060,7 +2060,7 @@ def _complete_ragflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes return response -def _complete_xai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_xai(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -2111,7 +2111,7 @@ def _complete_xai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: return response -def _complete_groq(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_groq(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -2174,7 +2174,7 @@ def _complete_groq(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult def _complete_bedrock_mantle( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -2219,7 +2219,7 @@ def _complete_bedrock_mantle( ) -def _complete_a2a(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_a2a(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -2282,7 +2282,7 @@ def _complete_a2a(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: ) -def _complete_gigachat(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_gigachat(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key = ctx.api_key @@ -2344,7 +2344,7 @@ def _complete_gigachat(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe return response -def _complete_sap(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_sap(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -2391,7 +2391,7 @@ def _complete_sap(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: def _complete_aiohttp_openai( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: if http2_enabled(): verbose_logger.warning( @@ -2451,7 +2451,7 @@ def _complete_aiohttp_openai( ) -def _complete_cometapi(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_cometapi(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -2500,7 +2500,7 @@ def _complete_cometapi(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe return response -def _complete_minimax(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_minimax(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -2546,7 +2546,7 @@ def _complete_minimax(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes return response -def _complete_hosted_vllm(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_hosted_vllm(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key: Final = ctx.api_key @@ -2591,7 +2591,7 @@ def _complete_hosted_vllm(ctx: _CompletionDispatchContext) -> _CompletionDispatc def _complete_custom_openai( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -2731,7 +2731,7 @@ def _complete_custom_openai( return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract -def _complete_mistral(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_mistral(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -2773,7 +2773,7 @@ def _complete_mistral(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes ) -def _complete_replicate(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_replicate(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key: Final = ctx.api_key @@ -2828,7 +2828,7 @@ def _complete_replicate(ctx: _CompletionDispatchContext) -> _CompletionDispatchR def _complete_anthropic_text( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -2882,7 +2882,7 @@ def _complete_anthropic_text( ) -def _complete_anthropic(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_anthropic(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -2948,7 +2948,7 @@ def _complete_anthropic(ctx: _CompletionDispatchContext) -> _CompletionDispatchR return response -def _complete_nlp_cloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_nlp_cloud(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key: Final = ctx.api_key @@ -2998,7 +2998,7 @@ def _complete_nlp_cloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchR return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract -def _complete_aleph_alpha(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_aleph_alpha(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: api_base = ctx.api_base api_key: Final = ctx.api_key litellm_params: Final = ctx.litellm_params @@ -3047,7 +3047,7 @@ def _complete_aleph_alpha(ctx: _CompletionDispatchContext) -> _CompletionDispatc return model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract -def _complete_cohere_chat(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_cohere_chat(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key: Final = ctx.api_key @@ -3112,7 +3112,7 @@ def _complete_cohere_chat(ctx: _CompletionDispatchContext) -> _CompletionDispatc ) -def _complete_maritalk(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_maritalk(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: api_base = ctx.api_base api_key: Final = ctx.api_key custom_prompt_dict: Final = ctx.custom_prompt_dict @@ -3145,7 +3145,7 @@ def _complete_maritalk(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe ) -def _complete_amazon_nova(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_amazon_nova(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: api_base = ctx.api_base api_key = ctx.api_key custom_llm_provider: Final = ctx.custom_llm_provider @@ -3181,7 +3181,7 @@ def _complete_amazon_nova(ctx: _CompletionDispatchContext) -> _CompletionDispatc ) -def _complete_huggingface(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_huggingface(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -3224,7 +3224,7 @@ def _complete_huggingface(ctx: _CompletionDispatchContext) -> _CompletionDispatc ) -def _complete_oci(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_oci(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -3259,7 +3259,7 @@ def _complete_oci(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: ) -def _complete_compactifai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_compactifai(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -3301,7 +3301,7 @@ def _complete_compactifai(ctx: _CompletionDispatchContext) -> _CompletionDispatc ) -def _complete_oobabooga(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_oobabooga(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: api_base: Final = ctx.api_base litellm_params: Final = ctx.litellm_params logger_fn: Final = ctx.logger_fn @@ -3335,7 +3335,7 @@ def _complete_oobabooga(ctx: _CompletionDispatchContext) -> _CompletionDispatchR return model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract -def _complete_databricks(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_databricks(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -3407,7 +3407,7 @@ def _complete_databricks(ctx: _CompletionDispatchContext) -> _CompletionDispatch return response -def _complete_datarobot(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_datarobot(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -3444,7 +3444,7 @@ def _complete_datarobot(ctx: _CompletionDispatchContext) -> _CompletionDispatchR ) -def _complete_openrouter(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_openrouter(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -3521,7 +3521,7 @@ def _complete_openrouter(ctx: _CompletionDispatchContext) -> _CompletionDispatch return response -def _complete_nadir(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_nadir(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: api_base: Final = ctx.api_base or litellm.api_base or get_secret_str("NADIR_API_BASE") or NADIR_DEFAULT_API_BASE api_key: Final = ctx.api_key @@ -3549,7 +3549,7 @@ def _complete_nadir(ctx: _CompletionDispatchContext) -> _CompletionDispatchResul def _complete_vercel_ai_gateway( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -3626,7 +3626,7 @@ def _complete_vercel_ai_gateway( return response -def _complete_edenai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_edenai(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: api_base: Final = litellm.EdenAIChatConfig.get_api_base(ctx.api_base) api_key: Final = litellm.EdenAIChatConfig.get_api_key(ctx.api_key or litellm.api_key) response: Final = base_llm_http_handler.completion( @@ -3652,7 +3652,7 @@ def _complete_edenai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu return response -def _complete_fal_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_fal_ai(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: if ctx.stream: raise litellm.FalAIError( status_code=400, @@ -3684,7 +3684,7 @@ def _complete_fal_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu def _complete_vertex_ai_beta( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -3751,7 +3751,7 @@ def _complete_vertex_ai_beta( ) -def _complete_vertex_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_vertex_ai(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base client: Final = _dispatch_client_http(ctx) @@ -3936,7 +3936,7 @@ def _complete_vertex_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchR return model_response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract -def _complete_predibase(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_predibase(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -3996,7 +3996,7 @@ def _complete_predibase(ctx: _CompletionDispatchContext) -> _CompletionDispatchR def _complete_text_completion_codestral( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -4046,7 +4046,7 @@ def _complete_text_completion_codestral( def _complete_text_completion_inception( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -4110,7 +4110,7 @@ def _complete_text_completion_inception( def _complete_sagemaker_chat( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base @@ -4146,7 +4146,7 @@ def _complete_sagemaker_chat( ) -def _complete_sagemaker(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_sagemaker(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion custom_prompt_dict: Final = ctx.custom_prompt_dict hf_model_name: Final = ctx.hf_model_name @@ -4178,7 +4178,7 @@ _ADDITIONAL_DROP_PARAMS_ADAPTER: Final = TypeAdapter(list[str]) _OPTIONAL_PARAMS_ADAPTER: Final = TypeAdapter(dict[str, object]) -def _complete_bedrock(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_bedrock(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -4307,7 +4307,7 @@ def _complete_bedrock(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes return response -def _complete_watsonx(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_watsonx(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -4345,7 +4345,7 @@ def _complete_watsonx(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes def _complete_watsonx_text( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base @@ -4421,7 +4421,7 @@ def _complete_watsonx_text( ) -def _complete_vllm(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_vllm(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: custom_prompt_dict = ctx.custom_prompt_dict litellm_params: Final = ctx.litellm_params logger_fn: Final = ctx.logger_fn @@ -4458,7 +4458,7 @@ def _complete_vllm(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult return model_response -def _complete_ollama(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_ollama(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key: Final = ctx.api_key @@ -4498,7 +4498,7 @@ def _complete_ollama(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu ) -def _complete_ollama_chat(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_ollama_chat(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -4540,7 +4540,7 @@ def _complete_ollama_chat(ctx: _CompletionDispatchContext) -> _CompletionDispatc ) -def _complete_triton(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_triton(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key: Final = ctx.api_key @@ -4576,7 +4576,7 @@ def _complete_triton(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu ) -def _complete_cloudflare(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_cloudflare(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -4615,7 +4615,7 @@ def _complete_cloudflare(ctx: _CompletionDispatchContext) -> _CompletionDispatch ) -def _complete_petals(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_petals(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: api_base = ctx.api_base client: Final = _dispatch_client_http(ctx) litellm_params: Final = ctx.litellm_params @@ -4655,7 +4655,7 @@ def _complete_petals(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu return model_response -def _complete_snowflake(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_snowflake(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key: Final = ctx.api_key @@ -4708,7 +4708,7 @@ def _complete_snowflake(ctx: _CompletionDispatchContext) -> _CompletionDispatchR return response -def _complete_gradient_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_gradient_ai(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key: Final = ctx.api_key @@ -4743,7 +4743,7 @@ def _complete_gradient_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatc ) -def _complete_gdc(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_gdc(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -4782,7 +4782,7 @@ def _complete_gdc(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: ) -def _complete_bytez(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_bytez(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key = ctx.api_key @@ -4822,7 +4822,7 @@ def _complete_bytez(ctx: _CompletionDispatchContext) -> _CompletionDispatchResul return response -def _complete_lemonade(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_lemonade(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base api_key = ctx.api_key @@ -4862,7 +4862,7 @@ def _complete_lemonade(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe return response -def _complete_ovhcloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_ovhcloud(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -4913,7 +4913,7 @@ def _custom_api_first_output(resp: httpx.Response | None) -> str: return resp.json()["data"][0]["output"][0] -def _complete_custom(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_custom(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: api_base: Final = ctx.api_base headers: Final = ctx.headers kwargs: Final = ctx.kwargs @@ -4983,7 +4983,7 @@ def _complete_custom(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu def _complete_custom_providers( - ctx: _CompletionDispatchContext, + ctx: CompletionDispatchContext, ) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base @@ -5045,7 +5045,7 @@ def _complete_custom_providers( return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract -def _complete_langgraph(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_langgraph(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -5094,7 +5094,7 @@ def _complete_langgraph(ctx: _CompletionDispatchContext) -> _CompletionDispatchR ) -def _complete_langflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: +def _complete_langflow(ctx: CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base = ctx.api_base api_key = ctx.api_key @@ -5277,7 +5277,7 @@ def completion( # Check if MCP tools are present (following responses pattern) # Cast tools to Optional[Iterable[ToolParam]] for type checking tools_for_mcp: Final = cast(Iterable[ToolParam] | None, tools) - if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools_for_mcp): + if LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(tools=tools_for_mcp): return acompletion_with_mcp( # pyright: ignore[reportReturnType] # MCP path returns a coroutine that acompletion() awaits; completion()'s sync return type omits it model=model, messages=messages, @@ -5659,7 +5659,9 @@ def completion( if litellm.add_function_to_prompt and optional_params.get( "functions_unsupported_model", None ): # if user opts to add it to prompt, when API doesn't support function calling - functions_unsupported_model: Final = optional_params.pop("functions_unsupported_model") + functions_unsupported_model: Final[list[object]] = TypeAdapter(list[object]).validate_python( + optional_params.pop("functions_unsupported_model") + ) messages = function_call_prompt(messages=messages, functions=functions_unsupported_model) # For logging - save the values of the litellm-specific params passed in @@ -5748,7 +5750,7 @@ def completion( raise litellm.BadRequestError( message=str(affinity_error), model=model, llm_provider=custom_llm_provider ) from affinity_error - cast(LiteLLMLoggingObj, logging).update_environment_variables( + logging.update_environment_variables( model=model, user=user, optional_params=processed_non_default_params, # [IMPORTANT] - using processed_non_default_params ensures consistent params logged to langfuse for finetuning / eval datasets. @@ -5840,7 +5842,7 @@ def completion( ): optional_params, _ = strip_reasoning_summary_aliases_from_optional_params(optional_params) - _dispatch_ctx: Final = _CompletionDispatchContext( + _dispatch_ctx: Final = CompletionDispatchContext( _azure_detection_model=_azure_detection_model, acompletion=acompletion, api_base=api_base, @@ -8772,7 +8774,7 @@ async def ahealth_check( log_raw_request_response=True, ) model_params["litellm_logging_obj"] = litellm_logging_obj - model_params = HealthCheckHelpers._update_model_params_with_health_check_tracking_information( + model_params = HealthCheckHelpers.update_model_params_with_health_check_tracking_information( model_params=model_params ) ######################################################### @@ -8816,11 +8818,11 @@ async def ahealth_check( _response: Final = await mode_handlers[mode]() # Only process headers for chat mode _response_headers: Final[dict] = getattr(_response, "_hidden_params", {}).get("headers", {}) or {} - return _create_health_check_response(_response_headers) + return create_health_check_response(_response_headers) else: raise Exception(f"Mode {mode} not supported. See modes here: https://docs.litellm.ai/docs/proxy/health") except Exception as e: - stack_trace = _redact_string(traceback.format_exc()) + stack_trace = redact_string(traceback.format_exc()) if isinstance(stack_trace, str): stack_trace = stack_trace[:1000] @@ -8967,7 +8969,7 @@ def _stamp_streaming_usage_cost(usage: Usage, response: ModelResponse, logging_o return if isinstance(getattr(usage, "cost", None), (int, float)): return - computed_cost: Final = logging_obj._response_cost_calculator(result=response) + computed_cost: Final = logging_obj.response_cost_calculator(result=response) if isinstance(computed_cost, (int, float)) and computed_cost > 0: setattr(usage, "cost", computed_cost) @@ -9098,7 +9100,7 @@ def stream_chunk_builder( if len(tool_call_chunks) > 0: tool_calls_list: Final = processor.get_combined_tool_content(tool_call_chunks) - _choice = cast(Choices, response.choices[0]) + _choice = response.choices[0] _choice.message.content = None _choice.message.tool_calls = tool_calls_list @@ -9111,7 +9113,7 @@ def stream_chunk_builder( ] if len(function_call_chunks) > 0: - _choice = cast(Choices, response.choices[0]) + _choice = response.choices[0] _choice.message.content = None _choice.message.function_call = processor.get_combined_function_call_content(function_call_chunks) @@ -9178,7 +9180,7 @@ def stream_chunk_builder( ] if len(audio_chunks) > 0: - _choice = cast(Choices, response.choices[0]) + _choice = response.choices[0] _choice.message.audio = processor.get_combined_audio_content(audio_chunks) # Handle image chunks from models like gemini-2.5-flash-image @@ -9229,7 +9231,7 @@ def stream_chunk_builder( } if combined_provider_fields: - _choice = cast(Choices, response.choices[0]) + _choice = response.choices[0] _choice.message.provider_specific_fields = combined_provider_fields completion_output = get_content_from_model_response(response) @@ -9396,9 +9398,9 @@ def _get_encoding() -> Tokenizer: def _load_default_encoding() -> Tokenizer: - from litellm._lazy_imports import _get_default_encoding + from litellm._lazy_imports import get_default_encoding - return _get_default_encoding() + return get_default_encoding() def __getattr__(name: str) -> Tokenizer: diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index 45ab52690bf..17b97dee37f 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -117,11 +117,11 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]): except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic try: await self._response.aread() - except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic + except Exception: # noqa: BLE001, S110 # Safe catch-all for cleanup logic pass try: await self._response.aclose() - except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic + except Exception: # noqa: BLE001, S110 # Safe catch-all for cleanup logic pass raise return self @@ -164,7 +164,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]): self._start_flush() try: await self._response.aclose() - except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic + except Exception: # noqa: BLE001, S110 # Safe catch-all for cleanup logic pass raise else: @@ -191,7 +191,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]): if self._initialized: await self._iterator.aclose() await self._response.aclose() - except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic + except Exception: # noqa: BLE001, S110 # Safe catch-all for cleanup logic pass @@ -239,7 +239,7 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]): self._start_flush() try: self._response.close() - except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic + except Exception: # noqa: BLE001, S110 # Safe catch-all for cleanup logic pass raise else: @@ -260,7 +260,7 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]): self._start_flush() try: self._response.close() - except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic + except Exception: # noqa: BLE001, S110 # Safe catch-all for cleanup logic pass @@ -375,7 +375,7 @@ async def allm_passthrough_route( provider=LlmProviders(resolved_custom_llm_provider), model=model, ) - except Exception: # noqa: BLE001 S110 + except Exception: # noqa: BLE001, S110 # provider config is optional # If we can't get provider config, pass None pass @@ -581,11 +581,11 @@ def llm_passthrough_route( except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic try: response.read() - except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic + except Exception: # noqa: BLE001, S110 # Safe catch-all for cleanup logic pass try: response.close() - except Exception: # noqa: BLE001 S110 # Safe catch-all for cleanup logic + except Exception: # noqa: BLE001, S110 # Safe catch-all for cleanup logic pass raise diff --git a/litellm/proxy/_experimental/mcp_server/mcp_debug.py b/litellm/proxy/_experimental/mcp_server/mcp_debug.py index 7b50a478297..4691584ecbd 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_debug.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_debug.py @@ -244,7 +244,7 @@ class MCPDebug: def mask_secret(value: str | None) -> str: if not value: return "(none)" - return MCPDebug._masker._mask_value(value) + return MCPDebug._masker.mask_value(value) @staticmethod def is_debug_enabled(headers: dict[str, str]) -> bool: diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 40476230879..8c9fdfcf67d 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -283,7 +283,7 @@ async def _a2a_sse_event_source( so the caller can relay them instead of breaking the stream. """ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client - from litellm.types.agents import _normalize_a2a_jsonrpc_response + from litellm.types.agents import normalize_a2a_jsonrpc_response from litellm.types.llms.custom_http import httpxSpecialProvider headers: Final = { @@ -305,7 +305,7 @@ async def _a2a_sse_event_source( try: parsed: Final = json.loads(error_body) if isinstance(parsed, dict) and "error" in parsed: - error_event = _normalize_a2a_jsonrpc_response(parsed, request_id=request_id) + error_event = normalize_a2a_jsonrpc_response(parsed, request_id=request_id) except Exception: error_event = None yield error_event or { @@ -867,7 +867,7 @@ async def invoke_agent_a2a( ) # Defer spend-log until after post_call_success_hook so guardrail # results written by the unified_guardrail hook are captured. - logging_obj._defer_async_logging = True + logging_obj.defer_async_logging = True response = await asend_message( model=f"a2a_agent/{agent_name}", request=a2a_request, @@ -887,9 +887,9 @@ async def invoke_agent_a2a( response=response, ) finally: - _enqueue_fn: Final = getattr(logging_obj, "_enqueue_deferred_logging", None) + _enqueue_fn: Final = getattr(logging_obj, "enqueue_deferred_logging", None) if _enqueue_fn is not None: - logging_obj._enqueue_deferred_logging = None + logging_obj.enqueue_deferred_logging = None _enqueue_fn() response_dict: Final[dict[str, object]] = ( diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 33fcac6230a..3848f5965bf 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -4588,7 +4588,7 @@ def _can_object_call_model( litellm.model_alias_map[model] if model in litellm.model_alias_map else ( - llm_router._get_model_from_alias(model) + llm_router.get_model_from_alias(model) if llm_router is not None and model in llm_router.model_group_alias else None ) diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 6445f0d7b05..8b2e0beb10f 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -228,7 +228,7 @@ class RouteChecks: # Use SensitiveDataMasker with custom configuration for user_id masker: Final = SensitiveDataMasker(visible_prefix=6, visible_suffix=2, mask_char="*") - return masker._mask_value(user_id) + return masker.mask_value(user_id) @staticmethod def _raise_admin_only_route_exception( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 687df9b0348..fcc50e57fd5 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2450,7 +2450,7 @@ class ProxyBaseLLMRequestProcessing: stored_cost: Final = logging_obj.model_call_details.get("response_cost") if isinstance(stored_cost, (int, float)): return float(stored_cost) - recomputed_cost: Final = logging_obj._response_cost_calculator(result=response) + recomputed_cost: Final = logging_obj.response_cost_calculator(result=response) return recomputed_cost if isinstance(recomputed_cost, (int, float)) else "" def _debug_log_request_payload(self) -> None: @@ -2609,7 +2609,7 @@ class ProxyBaseLLMRequestProcessing: if _post_call_guardrails_active and not self._is_streaming_request( data=self.data, is_streaming_request=is_streaming_request ): - logging_obj._defer_async_logging = True + logging_obj.defer_async_logging = True tasks: Final = [] # Start the moderation check (during_call_hook) as early as possible @@ -3172,7 +3172,7 @@ class ProxyBaseLLMRequestProcessing: except HTTPException: return - logging_obj._on_detached_stream_failure = _on_detached_stream_failure + logging_obj.on_detached_stream_failure = _on_detached_stream_failure def _is_streaming_response(self, response: object) -> bool: """ @@ -3428,10 +3428,10 @@ class ProxyBaseLLMRequestProcessing: if pending is not None: logging_obj._native_pending_logging = None # rebind-ok: consume the native OCR release signal once pending.release(not exception_raised) - _enqueue_fn: Final = getattr(logging_obj, "_enqueue_deferred_logging", None) + _enqueue_fn: Final = getattr(logging_obj, "enqueue_deferred_logging", None) if _enqueue_fn is None: return - logging_obj._enqueue_deferred_logging = None + logging_obj.enqueue_deferred_logging = None if exception_raised: return try: @@ -3478,7 +3478,7 @@ class ProxyBaseLLMRequestProcessing: from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.router_utils.add_retry_fallback_headers import HiddenParamsAsyncIteratorWrapper - unwrapped: Final = response._inner if isinstance(response, HiddenParamsAsyncIteratorWrapper) else response + unwrapped: Final = response.inner if isinstance(response, HiddenParamsAsyncIteratorWrapper) else response if isinstance(unwrapped, CustomStreamWrapper): # Intentionally a live reference (not a copy) — mirrors @@ -4223,7 +4223,7 @@ class ProxyBaseLLMRequestProcessing: debug_missing: Final = object() debug_before: Final = call_details.get(debug_key, debug_missing) if isinstance(call_details, dict) else None try: - cost: Final = litellm_logging_obj._response_cost_calculator(result=model_response) # pyright: ignore[reportPrivateUsage] # reuse the call's own cost calc for pricing parity with the logging callback + cost: Final = litellm_logging_obj.response_cost_calculator(result=model_response) except Exception: # noqa: BLE001 # a pricing failure falls back to model-name pricing instead of breaking the stream return None finally: diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index cb8b51d092e..9ab39fa9303 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -176,7 +176,7 @@ def initialize_callbacks_on_proxy( # check if callback is a custom logger compatible callback if isinstance(callback, str): - callback = LoggingCallbackManager._add_custom_callback_generic_api_str(callback) + callback = LoggingCallbackManager.add_custom_callback_generic_api_str(callback) if isinstance(callback, str) and callback in litellm._known_custom_logger_compatible_callbacks: imported_list.append(callback) elif isinstance(callback, str) and callback == "presidio": @@ -231,7 +231,7 @@ def initialize_callbacks_on_proxy( elif isinstance(callback, str) and callback == "openai_moderations": try: from enterprise.enterprise_hooks.openai_moderation import ( - _ENTERPRISE_OpenAI_Moderation, + ENTERPRISE_OpenAI_Moderation, ) except ImportError: raise Exception( @@ -242,7 +242,7 @@ def initialize_callbacks_on_proxy( if premium_user is not True: raise Exception("Trying to use OpenAI Moderations Check" + CommonProxyErrors.not_premium_user.value) - openai_moderations_object = _ENTERPRISE_OpenAI_Moderation() + openai_moderations_object = ENTERPRISE_OpenAI_Moderation() imported_list.append(openai_moderations_object) elif isinstance(callback, str) and callback == "lakera_prompt_injection": from litellm.proxy.guardrails.guardrail_hooks.lakera_ai import ( @@ -266,7 +266,7 @@ def initialize_callbacks_on_proxy( elif isinstance(callback, str) and callback == "google_text_moderation": try: from enterprise.enterprise_hooks.google_text_moderation import ( - _ENTERPRISE_GoogleTextModeration, + ENTERPRISE_GoogleTextModeration, ) except ImportError: raise Exception( @@ -277,7 +277,7 @@ def initialize_callbacks_on_proxy( if premium_user is not True: raise Exception("Trying to use Google Text Moderation" + CommonProxyErrors.not_premium_user.value) - google_text_moderation_obj = _ENTERPRISE_GoogleTextModeration() + google_text_moderation_obj = ENTERPRISE_GoogleTextModeration() imported_list.append(google_text_moderation_obj) elif isinstance(callback, str) and callback == "llmguard_moderations": try: @@ -295,7 +295,7 @@ def initialize_callbacks_on_proxy( elif isinstance(callback, str) and callback == "blocked_user_check": try: from enterprise.enterprise_hooks.blocked_user_list import ( - _ENTERPRISE_BlockedUserList, + ENTERPRISE_BlockedUserList, ) except ImportError: raise Exception( @@ -305,12 +305,12 @@ def initialize_callbacks_on_proxy( if premium_user is not True: raise Exception("Trying to use ENTERPRISE BlockedUser" + CommonProxyErrors.not_premium_user.value) - blocked_user_list = _ENTERPRISE_BlockedUserList(prisma_client=prisma_client) + blocked_user_list = ENTERPRISE_BlockedUserList(prisma_client=prisma_client) imported_list.append(blocked_user_list) elif isinstance(callback, str) and callback == "banned_keywords": try: from enterprise.enterprise_hooks.banned_keywords import ( - _ENTERPRISE_BannedKeywords, + ENTERPRISE_BannedKeywords, ) except ImportError: raise Exception( @@ -320,7 +320,7 @@ def initialize_callbacks_on_proxy( if premium_user is not True: raise Exception("Trying to use ENTERPRISE BannedKeyword" + CommonProxyErrors.not_premium_user.value) - banned_keywords_obj = _ENTERPRISE_BannedKeywords() + banned_keywords_obj = ENTERPRISE_BannedKeywords() imported_list.append(banned_keywords_obj) elif isinstance(callback, str) and callback == "detect_prompt_injection": from litellm.proxy.hooks.prompt_injection_detection import ( diff --git a/litellm/proxy/common_utils/prompt_cache_pricing.py b/litellm/proxy/common_utils/prompt_cache_pricing.py index 95b945714e0..b07873db510 100644 --- a/litellm/proxy/common_utils/prompt_cache_pricing.py +++ b/litellm/proxy/common_utils/prompt_cache_pricing.py @@ -7,8 +7,8 @@ from pydantic import TypeAdapter import litellm from litellm.cost_calculator import ( - _select_model_name_for_cost_calc, # pyright: ignore[reportPrivateUsage] # shares completion_cost's deployment tariff selection completion_cost, # pyright: ignore[reportUnknownVariableType] # legacy optional parameters are untyped + select_model_name_for_cost_calc, ) from litellm.litellm_core_utils.litellm_logging import Logging from litellm.types.management_endpoints.prompt_cache_prediction import CacheTokenBuckets @@ -42,7 +42,7 @@ def price_cache_tokens( model: str, deployment_id: str, tokens: CacheTokenBuckets, completion_tokens: int = 0 ) -> float | None: try: - selected_model: Final = _select_model_name_for_cost_calc( + selected_model: Final = select_model_name_for_cost_calc( model=model, completion_response=None, custom_pricing=True, diff --git a/litellm/proxy/container_endpoints/ownership.py b/litellm/proxy/container_endpoints/ownership.py index 14480232a4a..aeb0bb1e3d0 100644 --- a/litellm/proxy/container_endpoints/ownership.py +++ b/litellm/proxy/container_endpoints/ownership.py @@ -67,7 +67,7 @@ def _container_model_object_id(original_container_id: str, custom_llm_provider: def decode_container_id_for_ownership(container_id: str, custom_llm_provider: str) -> tuple[str, str]: - decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) + decoded: Final = ResponsesAPIRequestUtils.decode_container_id(container_id) original_container_id: Final = decoded.get("response_id", container_id) decoded_provider: Final = decoded.get("custom_llm_provider") if decoded_provider and custom_llm_provider == "openai": @@ -82,7 +82,7 @@ async def get_container_forwarding_params( "container_id": original_container_id, "custom_llm_provider": custom_llm_provider, } - decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) + decoded: Final = ResponsesAPIRequestUtils.decode_container_id(container_id) model_id = decoded.get("model_id") if not (isinstance(model_id, str) and model_id): # Native upstream IDs (e.g. Azure ``cntr_``) carry no LiteLLM @@ -92,7 +92,7 @@ async def get_container_forwarding_params( # selected a specific deployment that ID embeds the model_id. stored_id: Final = await _get_stored_container_id(original_container_id, custom_llm_provider) if stored_id and stored_id != container_id: - stored_decoded: Final = ResponsesAPIRequestUtils._decode_container_id(stored_id) + stored_decoded: Final = ResponsesAPIRequestUtils.decode_container_id(stored_id) stored_model_id: Final = stored_decoded.get("model_id") if isinstance(stored_model_id, str) and stored_model_id: model_id = stored_model_id diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index af2b7f0df8a..6d25307db96 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -15,7 +15,7 @@ from pydantic import TypeAdapter import litellm from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.credential_accessor import CredentialAccessor -from litellm.litellm_core_utils.litellm_logging import _get_masked_values +from litellm.litellm_core_utils.litellm_logging import get_masked_values from litellm.llms.anthropic.wif import ( ExportedJwks, NotAnInternalIssuerCredential, @@ -269,7 +269,7 @@ async def get_credentials( masked_credentials: Final = [ { "credential_name": credential.credential_name, - "credential_values": _get_masked_values(credential.credential_values), + "credential_values": get_masked_values(credential.credential_values), "credential_info": credential.credential_info, } for credential in litellm.credential_list @@ -299,7 +299,7 @@ async def get_credential_by_name( if credential.credential_name == credential_name: masked_credential = CredentialItem( credential_name=credential.credential_name, - credential_values=_get_masked_values( + credential_values=get_masked_values( credential.credential_values, unmasked_length=4, number_of_asterisks=4, @@ -401,7 +401,7 @@ async def get_credential_by_model( credential_values: Final = llm_router.get_deployment_credentials(model_id) if credential_values is None: raise HTTPException(status_code=404, detail="Model not found") - masked_credential_values: Final = _get_masked_values( + masked_credential_values: Final = get_masked_values( credential_values, unmasked_length=4, number_of_asterisks=4, diff --git a/litellm/proxy/guardrails/auto_router_compression.py b/litellm/proxy/guardrails/auto_router_compression.py index f132b72e922..ecb182ab33c 100644 --- a/litellm/proxy/guardrails/auto_router_compression.py +++ b/litellm/proxy/guardrails/auto_router_compression.py @@ -157,14 +157,14 @@ async def arm_pre_call( return from litellm.router_strategy.tag_based_routing import ( - _get_tags_from_request_kwargs, # pyright: ignore[reportPrivateUsage] # used in router.py and budget_limiter.py too + get_tags_from_request_kwargs, ) policy: Final = policy_for_model( llm_router=llm_router, model_alias=model_alias, request_kwargs=data, - request_tags=_get_tags_from_request_kwargs(data), + request_tags=get_tags_from_request_kwargs(data), ) if policy is None: return diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index c43c59f11b8..1dcc7cf5485 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -127,12 +127,12 @@ def _get_guardrails_list_response( """ Helper function to get the guardrails list response """ - from litellm.litellm_core_utils.litellm_logging import _get_masked_values + from litellm.litellm_core_utils.litellm_logging import get_masked_values guardrail_configs: Final[list[GuardrailInfoResponse]] = [] for guardrail in guardrails_config: litellm_params = guardrail.get("litellm_params") or {} - masked_params = _get_masked_values( + masked_params = get_masked_values( litellm_params, unmasked_length=4, number_of_asterisks=4, @@ -241,7 +241,7 @@ async def list_guardrails_v2( } ``` """ - from litellm.litellm_core_utils.litellm_logging import _get_masked_values + from litellm.litellm_core_utils.litellm_logging import get_masked_values from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER from litellm.proxy.proxy_server import prisma_client @@ -277,13 +277,13 @@ async def list_guardrails_v2( if isinstance(litellm_params, LitellmParams) else litellm_params ) or {} - masked_litellm_params_dict = _get_masked_values( + masked_litellm_params_dict = get_masked_values( litellm_params_dict, unmasked_length=4, number_of_asterisks=4, ) masked_litellm_params = ( - BaseLitellmParams(**masked_litellm_params_dict) if masked_litellm_params_dict else None + BaseLitellmParams.model_validate(masked_litellm_params_dict) if masked_litellm_params_dict else None ) guardrail_configs.append( GuardrailInfoResponse( @@ -318,13 +318,15 @@ async def list_guardrails_v2( if isinstance(in_memory_litellm_params_raw, LitellmParams) else in_memory_litellm_params_raw ) or {} - masked_in_memory_litellm_params = _get_masked_values( + masked_in_memory_litellm_params = get_masked_values( in_memory_litellm_params_dict, unmasked_length=4, number_of_asterisks=4, ) masked_in_memory_litellm_params_typed = ( - BaseLitellmParams(**masked_in_memory_litellm_params) if masked_in_memory_litellm_params else None + BaseLitellmParams.model_validate(masked_in_memory_litellm_params) + if masked_in_memory_litellm_params + else None ) guardrail_configs.append( GuardrailInfoResponse( @@ -866,12 +868,12 @@ async def _get_user_team_ids(user_api_key_dict: UserAPIKeyAuth) -> list[str]: def _row_to_submission_item(row: "LiteLLM_GuardrailsTable") -> GuardrailSubmissionItem: - from litellm.litellm_core_utils.litellm_logging import _get_masked_values + from litellm.litellm_core_utils.litellm_logging import get_masked_values guardrail_info: Final = _parse_json_field(row.guardrail_info) or {} team_guardrail: Final = row.team_id is not None raw_params: Final = decrypt_guardrail_litellm_params(_parse_json_field(row.litellm_params) or {}) - masked_params: Final = _get_masked_values(raw_params, unmasked_length=4, number_of_asterisks=4) + masked_params: Final = get_masked_values(raw_params, unmasked_length=4, number_of_asterisks=4) return GuardrailSubmissionItem( guardrail_id=row.guardrail_id, guardrail_name=row.guardrail_name, @@ -1366,7 +1368,7 @@ async def get_guardrail_info(guardrail_id: str): ``` """ - from litellm.litellm_core_utils.litellm_logging import _get_masked_values + from litellm.litellm_core_utils.litellm_logging import get_masked_values from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER from litellm.proxy.proxy_server import prisma_client from litellm.types.guardrails import GUARDRAIL_DEFINITION_LOCATION @@ -1396,12 +1398,14 @@ async def get_guardrail_info(guardrail_id: str): if isinstance(litellm_params, LitellmParams) else litellm_params ) or {} - masked_litellm_params_dict: Final = _get_masked_values( + masked_litellm_params_dict: Final = get_masked_values( result_litellm_params_dict, unmasked_length=4, number_of_asterisks=4, ) - masked_litellm_params = BaseLitellmParams(**masked_litellm_params_dict) if masked_litellm_params_dict else None + masked_litellm_params = ( + BaseLitellmParams.model_validate(masked_litellm_params_dict) if masked_litellm_params_dict else None + ) return GuardrailInfoResponse( guardrail_id=result.get("guardrail_id"), diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index c5104254c4e..37318feba45 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -35,7 +35,7 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys from litellm.litellm_core_utils.litellm_logging import ( - _get_masked_values, # pyright: ignore[reportPrivateUsage] # the shared header-masking helper has no public name + get_masked_values, ) from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import ( bedrock_guardrail_cost_by_unit, @@ -1233,7 +1233,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): "Bedrock AI request body: %s, url %s, headers: %s", bedrock_request_data, prepared_request.url, - _get_masked_values(headers_dict), + get_masked_values(headers_dict), ) httpx_response: Final = await self._sign_and_post( diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 0c4ea6b29b5..03796569043 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -18,7 +18,7 @@ from fastapi import HTTPException from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache -from litellm.cost_calculator import _infer_call_type +from litellm.cost_calculator import infer_call_type from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route @@ -96,7 +96,7 @@ def resolve_endpoint_translation( route_call_types[0].value if route_call_types else ( - _infer_call_type(call_type=None, completion_response=first_response_item) + infer_call_type(call_type=None, completion_response=first_response_item) if first_response_item is not None else None ) @@ -340,7 +340,7 @@ class UnifiedLLMGuardrails(CustomLogger): if call_types is not None and len(call_types) > 0: call_type = call_types[0] if call_type is None: - call_type = _infer_call_type(call_type=None, completion_response=response) + call_type = infer_call_type(call_type=None, completion_response=response) if call_type is None: litellm_logging_obj: Final = data.get("litellm_logging_obj") @@ -1213,7 +1213,7 @@ class UnifiedLLMGuardrails(CustomLogger): call_type = call_types[0].value if call_type is None: - call_type = _infer_call_type(call_type=None, completion_response=item) + call_type = infer_call_type(call_type=None, completion_response=item) # If call type not supported, just pass through all chunks if call_type is None or CallTypes(call_type) not in mappings: diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 4c077b02c00..1ad01d19e50 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -998,7 +998,7 @@ def _caller_may_probe_deployment( caller_is_admin: bool, ) -> bool: """Same deployment visibility rule as routing: another team's deployment is never in scope, team-less callers included.""" - if not caller_is_admin and not Router._deployment_usable_by_team(deployment, team_id): + if not caller_is_admin and not Router.deployment_usable_by_team(deployment, team_id): return False if allowed_models is None: return True diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 9f5bc0966f8..55bd622651e 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -30,10 +30,10 @@ import litellm from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger from litellm.batches.batch_utils import ( - _count_entry_tokens, - _estimate_batch_entry_tokens, - _extract_file_access_credentials, - _iter_batch_input_lines, + count_entry_tokens, + estimate_batch_entry_tokens, + extract_file_access_credentials, + iter_batch_input_lines, ) from litellm.constants import BATCH_TPD_DESCRIPTOR_SUFFIX, BATCH_TPD_WINDOW_SECONDS from litellm.exceptions import RateLimitErrorCategory @@ -537,7 +537,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): model_id=model_from_file_id, operation_context="batch input file read (rate limiting)", ) - fetch_kwargs.update(_extract_file_access_credentials(credentials)) + fetch_kwargs.update(extract_file_access_credentials(credentials)) fetch_kwargs["model"] = model_from_file_id provider = credentials.get("custom_llm_provider") if provider: @@ -554,7 +554,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): model_id=request_model, operation_context="batch input file read (rate limiting)", ) - fetch_kwargs.update(_extract_file_access_credentials(credentials)) + fetch_kwargs.update(extract_file_access_credentials(credentials)) fetch_kwargs["model"] = request_model provider = credentials.get("custom_llm_provider") if provider: @@ -947,12 +947,12 @@ class _PROXY_BatchRateLimiter(CustomLogger): total_tokens = 0 output_tokens = 0 # rebind-ok: accumulated per JSONL row in the loop below request_count = 0 - for raw_line in _iter_batch_input_lines(file_content_bytes): + for raw_line in iter_batch_input_lines(file_content_bytes): request_count += 1 try: entry = json.loads(raw_line) except Exception: - entry_total_tokens = _estimate_batch_entry_tokens(raw_line) + entry_total_tokens = estimate_batch_entry_tokens(raw_line) entry_output_tokens = self.parallel_request_limiter.no_max_tokens_output_floor( min_configured_otpm_limit ) @@ -973,9 +973,9 @@ class _PROXY_BatchRateLimiter(CustomLogger): output_tokens += entry_output_tokens try: - entry_total_tokens = _count_entry_tokens(entry) + entry_total_tokens = count_entry_tokens(entry) except Exception: - entry_total_tokens = _estimate_batch_entry_tokens(raw_line) + entry_total_tokens = estimate_batch_entry_tokens(raw_line) total_tokens += entry_total_tokens if model: diff --git a/litellm/proxy/hooks/batch_redis_get.py b/litellm/proxy/hooks/batch_redis_get.py index 3d4ef67bb91..94ad6fb7cb2 100644 --- a/litellm/proxy/hooks/batch_redis_get.py +++ b/litellm/proxy/hooks/batch_redis_get.py @@ -110,6 +110,6 @@ class _PROXY_BatchRedisRequests(CustomLogger): cached_result = await litellm.cache.cache.async_get_cache(cache_key, *args, **kwargs) if cached_result is not None: await self.in_memory_cache.async_set_cache(cache_key, cached_result, ttl=60) - return litellm.cache._get_cache_logic(cached_result=cached_result, max_age=max_age) + return litellm.cache.get_cache_logic(cached_result=cached_result, max_age=max_age) except Exception: return None diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index d600e249754..0a078321d5e 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -698,7 +698,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): - priority_model: Priority-specific token tracking """ from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, + get_parent_otel_span_from_kwargs, ) from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, @@ -709,7 +709,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): try: verbose_proxy_logger.debug("INSIDE dynamic rate limiter ASYNC SUCCESS LOGGING") - litellm_parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + litellm_parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) # Get metadata from standard_logging_object standard_logging_object: Final = kwargs.get("standard_logging_object") or {} diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py index bec12ba8201..15c28d64d05 100644 --- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py +++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py @@ -86,7 +86,7 @@ class SemanticToolFilterHook(CustomLogger): LiteLLM_Proxy_MCP_Handler, ) - return LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools) + return LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(tools) async def _expand_mcp_tools( self, @@ -104,7 +104,7 @@ class SemanticToolFilterHook(CustomLogger): ) # Parse to separate MCP tools from other tools - mcp_tools, _ = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools) + mcp_tools, _ = await LiteLLM_Proxy_MCP_Handler.split_mcp_tools(tools) if not mcp_tools: return [] @@ -114,7 +114,7 @@ class SemanticToolFilterHook(CustomLogger): ( openai_tools, _, - ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_to_openai_format( + ) = await LiteLLM_Proxy_MCP_Handler.process_mcp_tools_to_openai_format( user_api_key_auth=user_api_key_dict, mcp_tools_with_litellm_proxy=mcp_tools ) diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index b4ce010dd27..facae4cbb81 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -12,7 +12,7 @@ from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger from litellm.exceptions import RateLimitType from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs +from litellm.litellm_core_utils.core_helpers import get_parent_otel_span_from_kwargs from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( get_key_model_rpm_limit, @@ -499,7 +499,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): get_model_group_from_litellm_kwargs, ) - litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs) + litellm_parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(kwargs=kwargs) try: self.print_verbose("INSIDE parallel request limiter ASYNC SUCCESS LOGGING") @@ -703,7 +703,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): try: self.print_verbose("Inside Max Parallel Request Failure Hook") - litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs) + litellm_parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(kwargs=kwargs) _metadata: Final = kwargs["litellm_params"].get("metadata", {}) or {} global_max_parallel_requests: Final = _metadata.get("global_max_parallel_requests", None) user_api_key: Final = _metadata.get("user_api_key", None) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 363d525baf3..882c4d126cf 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -5002,12 +5002,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): Update TPM usage on successful API calls by incrementing counters using pipeline """ from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, + get_parent_otel_span_from_kwargs, ) rate_limit_type: Final = self.get_rate_limit_type() - litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs) + litellm_parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(kwargs) try: verbose_proxy_logger.debug("INSIDE parallel request limiter ASYNC SUCCESS LOGGING") @@ -5125,11 +5125,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): usage instead of refunding it. """ from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, + get_parent_otel_span_from_kwargs, ) try: - litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs) + litellm_parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(kwargs) pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = [] diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index e7160436638..ffe752329d4 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -10,10 +10,10 @@ from litellm.batches.batch_utils import batch_cost_is_final from litellm.constants import BACKGROUND_INTERACTION_COST_POLLING_ENABLED from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, budget_reservation_from_metadata, get_litellm_metadata_from_kwargs, get_metadata_variable_name_from_kwargs, + get_parent_otel_span_from_kwargs, ) from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost @@ -309,7 +309,7 @@ class _ProxyDBLogger(CustomLogger): kwargs.get("stream", None), kwargs.get("complete_streaming_response", None), ) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs=kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs=kwargs) litellm_params: Final = kwargs.get("litellm_params", {}) or {} end_user_id: Final = get_end_user_id_for_cost_tracking(litellm_params) metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index 77e5b0e5674..6b4f2b690db 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -185,7 +185,7 @@ def _mask_sensitive_fields(data: Mapping[str, object], sensitive_fields: set[str masked: Final[dict[str, object]] = {} for key, value in data.items(): if value is not None and key in sensitive_fields and isinstance(value, str): - masked[key] = _sensitive_masker._mask_value(value) + masked[key] = _sensitive_masker.mask_value(value) else: masked[key] = value return masked diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index 55352ede1dc..650e743b027 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -230,7 +230,7 @@ async def unblock_user(data: BlockUsers): """ try: from enterprise.enterprise_hooks.blocked_user_list import ( - _ENTERPRISE_BlockedUserList, + ENTERPRISE_BlockedUserList, ) except ImportError: raise HTTPException( @@ -242,7 +242,7 @@ async def unblock_user(data: BlockUsers): ) if ( - not any(isinstance(x, _ENTERPRISE_BlockedUserList) for x in litellm.callbacks) + not any(isinstance(x, ENTERPRISE_BlockedUserList) for x in litellm.callbacks) or litellm.blocked_user_list is None ): raise HTTPException( diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index f2dc2995e32..72844de2549 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -4509,7 +4509,7 @@ def _check_model_access_group(models: list[str] | None, llm_router: Router | Non return True for model in models: - if llm_router._is_model_access_group_for_wildcard_route(model_access_group=model): + if llm_router.is_model_access_group_for_wildcard_route(model_access_group=model): if not premium_user: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 3c1b31b68ef..e6705730e6f 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -76,7 +76,7 @@ from litellm.repositories.verification_token_repository import ( from litellm.types.proxy.management_endpoints.common_daily_activity import ( SpendAnalyticsPaginatedResponse, ) -from litellm.utils import _update_dictionary +from litellm.utils import update_dictionary if TYPE_CHECKING: from types import TracebackType @@ -770,7 +770,7 @@ async def update_organization( updated_metadata: Final = _STR_OBJECT_DICT_ADAPTER.validate_python( updated_organization_row_json.get("metadata", {}) ) - merged_metadata: Final[Mapping[str, object]] = _update_dictionary( + merged_metadata: Final[Mapping[str, object]] = update_dictionary( existing_dict=cast( # cast-ok: prisma de-serializes a Json column to the plain python dict it stores "dict[str, object]", existing_metadata ).copy(), diff --git a/litellm/proxy/management_endpoints/router_settings_endpoints.py b/litellm/proxy/management_endpoints/router_settings_endpoints.py index 3f927906ed8..aa483eb2f6e 100644 --- a/litellm/proxy/management_endpoints/router_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/router_settings_endpoints.py @@ -123,7 +123,7 @@ async def get_router_settings( # generic `hasattr` loop below would miss them. current_values["routing_groups"] = [ group.model_dump(exclude=frozenset(("model_priorities",)) if group.model_priorities is None else None) - for group in llm_router._routing_groups.values() + for group in llm_router.routing_groups.values() ] for field in router_fields: if field.field_name == "routing_groups": diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 4840b28cb39..7e7ab1dc1f0 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -104,7 +104,7 @@ class ManagedFileIdResolver(Protocol): ) -> Mapping[str, str]: ... -def _is_base64_encoded_unified_file_id(b64_uid: str) -> str | Literal[False]: +def _is_base64_encoded_unified_file_id(b64_uid: object) -> str | Literal[False]: # Ensure b64_uid is a string and not a mock object if not isinstance(b64_uid, str): return False @@ -1196,7 +1196,7 @@ def _model_name_for_batch_response(response: "LiteLLMBatch") -> str | None: ) -def _batch_owner_auth_from_db_object(db_batch_object: "LiteLLM_ManagedObjectTable") -> "UserAPIKeyAuth | None": +def _batch_owner_auth_from_db_object(db_batch_object: object) -> "UserAPIKeyAuth | None": from litellm.proxy._types import UserAPIKeyAuth created_by: Final = getattr(db_batch_object, "created_by", None) @@ -1284,7 +1284,7 @@ async def ensure_batch_response_managed_file_ids( prisma_client, verbose_proxy_logger, user_api_key_dict=None, - db_batch_object: "LiteLLM_ManagedObjectTable | None" = None, + db_batch_object: object | None = None, unified_batch_id: str | Literal[False] | None = None, ) -> None: """Normalize batch file IDs to managed unified IDs before DB persistence.""" diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index b712065b345..88c8a67568b 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -1864,13 +1864,13 @@ def _resolve_vertex_model_from_router( if not deployment: return encoded_endpoint, endpoint, vertex_project, vertex_location, None - litellm_params: Final = deployment.get("litellm_params", {}) + litellm_params: Final[Mapping[str, object]] = cast(Mapping[str, object], deployment.get("litellm_params", {})) model_info: Final = deployment.get("model_info") deployment_model_info: Final = model_info if isinstance(model_info, Mapping) else None # Always override with router config values (they take precedence over URL values) - config_vertex_project: Final = litellm_params.get("vertex_project") - config_vertex_location: Final = litellm_params.get("vertex_location") + config_vertex_project: Final = cast(str | None, litellm_params.get("vertex_project")) + config_vertex_location: Final = cast(str | None, litellm_params.get("vertex_location")) if config_vertex_project: vertex_project = config_vertex_project if config_vertex_location: @@ -1878,7 +1878,7 @@ def _resolve_vertex_model_from_router( # Get the actual Vertex AI model name by stripping the provider prefix # e.g., "vertex_ai/gemini-2.0-flash-exp" -> "gemini-2.0-flash-exp" - model_from_config: Final = litellm_params.get("model", "") + model_from_config: Final = cast(str, litellm_params.get("model", "")) if model_from_config: # get_llm_provider returns (model, custom_llm_provider, dynamic_api_key, api_base) # For "vertex_ai/gemini-2.0-flash-exp" it returns: @@ -2340,7 +2340,7 @@ async def azure_proxy_route( base_target_url=base_target_url, api_key=None, custom_llm_provider=litellm.LlmProviders.AZURE_AI, - extra_headers=cast(dict, extra_headers), + extra_headers=extra_headers, ) body_model_group_relay: Final = await _relay_azure_body_model_group( diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py index 0ff02b29c58..1e4a6e679f7 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py @@ -307,7 +307,7 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): usage: Final = self._session_usage(turn, model) if usage is None: return None - cost: Final = logging_obj._response_cost_calculator( # pyright: ignore[reportPrivateUsage] # the call's own calculator keeps custom pricing and the deployment's region in step with the spend row + cost: Final = logging_obj.response_cost_calculator( result=ModelResponse(model=model, usage=usage), litellm_model_name=model, ) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index f5ba7f9e877..3f4a30e9179 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -58,7 +58,7 @@ from litellm.litellm_core_utils.core_helpers import ( from litellm.litellm_core_utils.initialize_dynamic_callback_params import validate_no_callback_env_reference from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.litellm_core_utils.litellm_logging import _get_masked_values +from litellm.litellm_core_utils.litellm_logging import get_masked_values from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.litellm_core_utils.redact_messages import should_redact_message_logging from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @@ -1159,7 +1159,7 @@ async def pass_through_request( verbose_proxy_logger.debug( "Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n", url, - _get_masked_values(upstream_headers), + get_masked_values(upstream_headers), _parsed_body, ) diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index 19d8b063dd7..9a259b4778a 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -64,7 +64,7 @@ class PassThroughStreamingHandler: @staticmethod def _stamp_first_chunk_if_needed(litellm_logging_obj: LiteLLMLoggingObj) -> None: if litellm_logging_obj.completion_start_time is None: - litellm_logging_obj._update_completion_start_time(completion_start_time=datetime.now()) + litellm_logging_obj.update_completion_start_time(completion_start_time=datetime.now()) @staticmethod async def schedule_stream_failure_logging( diff --git a/litellm/proxy/prompts/prompt_endpoints.py b/litellm/proxy/prompts/prompt_endpoints.py index 41c8a6a77b4..b7775b99419 100644 --- a/litellm/proxy/prompts/prompt_endpoints.py +++ b/litellm/proxy/prompts/prompt_endpoints.py @@ -430,15 +430,15 @@ def _get_prompt_template(prompt_spec: PromptSpec, base_prompt_id: str) -> Prompt dotprompt_content: Final = prompt_spec.litellm_params.dotprompt_content if dotprompt_content: from litellm.integrations.dotprompt import ( - _get_prompt_data_from_dotprompt_content, + get_prompt_data_from_dotprompt_content, ) - parsed: Final = _get_prompt_data_from_dotprompt_content(dotprompt_content) + parsed: Final = get_prompt_data_from_dotprompt_content(dotprompt_content) if parsed: return PromptTemplateBase( litellm_prompt_id=base_prompt_id, - content=parsed.get("content", ""), - metadata=parsed.get("metadata"), + content=cast(str, parsed.get("content", "")), + metadata=cast(dict[str, object] | None, parsed.get("metadata")), ) else: prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec) @@ -1064,7 +1064,7 @@ async def test_prompt( try: # Parse the dotprompt content and create PromptTemplate prompt_manager: Final = PromptManager() - frontmatter, template_content = prompt_manager._parse_frontmatter(content=request.dotprompt_content) + frontmatter, template_content = prompt_manager.parse_frontmatter(content=request.dotprompt_content) # Create PromptTemplate to leverage existing parameter extraction logic template: Final = PromptTemplate(content=template_content, metadata=frontmatter, template_id="test_prompt") diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index faf268d3986..53a447bc6ec 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -1229,7 +1229,7 @@ def run_server( litellm.json_logs = True - litellm._turn_on_json() + litellm.turn_on_json() ### GENERAL SETTINGS ### general_settings = _config.get("general_settings", {}) if general_settings is None: @@ -1499,7 +1499,7 @@ def run_server( import litellm if detailed_debug is True: - litellm._turn_on_debug() + litellm.turn_on_debug() # DO NOT DELETE - enables global variables to work across files from litellm.proxy.proxy_server import app diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1f3f6ea0adb..a620d9320fa 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -269,7 +269,7 @@ import litellm import litellm._redis from litellm import Router from litellm._internal_context import service_target, with_service_target -from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger +from litellm._logging import redact_string, verbose_proxy_logger, verbose_router_logger from litellm.caching.caching import DualCache, RedisCache from litellm.caching.dual_cache import DeclaredBatchRead from litellm.caching.redis_batch import ( @@ -317,9 +317,9 @@ from litellm.litellm_core_utils.agentic_loop_settings import ( ) from litellm.litellm_core_utils.audio_utils.utils import resolve_speech_media_type from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, drop_params_flag, get_litellm_metadata_from_kwargs, + get_parent_otel_span_from_kwargs, ) from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -923,7 +923,7 @@ from litellm.types.secret_managers.main import ( ) from litellm.types.utils import CredentialItem, CustomHuggingfaceTokenizer, RawRequestTypedDict, StandardLoggingPayload from litellm.types.utils import ModelInfo as ModelMapInfo -from litellm.utils import _add_custom_logger_callback_to_specific_event +from litellm.utils import add_custom_logger_callback_to_specific_event try: from litellm._version import version @@ -1574,7 +1574,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState for _tagged_routers in llm_router.adaptive_routers.values(): for _tagged in _tagged_routers: await _tagged.strategy.load_state_from_db(prisma_client) - _tagged.strategy._state_loaded = True + _tagged.strategy.state_loaded = True asyncio.create_task(_adaptive_router_flusher_loop()) ## [Optional] Initialize dd tracer @@ -4457,7 +4457,7 @@ def _write_health_state_to_router_cache( is on. """ from litellm.proxy.health_check import build_deployment_health_states - from litellm.router_utils.cooldown_handlers import _set_cooldown_deployments + from litellm.router_utils.cooldown_handlers import set_cooldown_deployments from litellm.router_utils.router_callbacks.track_deployment_metrics import ( increment_deployment_failures_for_current_minute, ) @@ -4516,7 +4516,7 @@ def _write_health_state_to_router_cache( deployment_id=model_id, ) - _set_cooldown_deployments( + set_cooldown_deployments( litellm_router_instance=llm_router, original_exception=original_exception, exception_status=exception_status, @@ -4549,11 +4549,11 @@ async def _adaptive_router_flusher_loop(): ar = tagged.strategy # Lazy state load: covers adaptive routers registered via # `/config/reload` after proxy boot. - if not getattr(ar, "_state_loaded", False): + if not getattr(ar, "state_loaded", False): try: await ar.load_state_from_db(prisma_client) finally: - ar._state_loaded = True + ar.state_loaded = True await ar.queue.flush_state_to_db(prisma_client) await ar.queue.flush_session_to_db(prisma_client) except asyncio.CancelledError: @@ -6666,7 +6666,7 @@ class ProxyConfig: litellm.key_alias_pattern = parse_key_alias_pattern(value) elif key == "json_logs" and value is True: litellm.json_logs = True - litellm._turn_on_json() + litellm.turn_on_json() verbose_proxy_logger.debug( "%s Enabled JSON logging via config%s", blue_color_code, reset_color_code ) @@ -7076,7 +7076,7 @@ class ProxyConfig: ) if redis_usage_cache is not None and router.cache.redis_cache is None: - router._update_redis_cache(cache=redis_usage_cache) + router.update_redis_cache(cache=redis_usage_cache) # Guardrail settings guardrails_v2: list[dict] | None = None @@ -7610,7 +7610,7 @@ class ProxyConfig: """ if callback in litellm._known_custom_logger_compatible_callbacks: for event_type in event_types: - _add_custom_logger_callback_to_specific_event(callback, event_type) + add_custom_logger_callback_to_specific_event(callback, event_type) elif callback not in existing_callbacks: if event_types == ["success"]: litellm.logging_callback_manager.add_litellm_success_callback(callback) @@ -8954,7 +8954,7 @@ class ProxyConfig: try: # read vector stores from db table - vector_stores: Final = await VectorStoreRegistry._get_vector_stores_from_db(prisma_client=prisma_client) + vector_stores: Final = await VectorStoreRegistry.get_vector_stores_from_db(prisma_client=prisma_client) if len(vector_stores) <= 0: return @@ -8973,7 +8973,7 @@ class ProxyConfig: try: # read vector stores from db table - vector_store_indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db( + vector_store_indexes: Final = await VectorStoreIndexRegistry.get_vector_store_indexes_from_db( prisma_client=prisma_client ) @@ -10383,7 +10383,7 @@ class ProxyStartupEvent: enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, ) if llm_router is not None and llm_router.cache.redis_cache is None: - llm_router._update_redis_cache(cache=coordination_redis_cache) + llm_router.update_redis_cache(cache=coordination_redis_cache) verbose_proxy_logger.info( "coordination_redis: using the standalone Redis saved in the database " "for usage tracking, rate limiting, and cross-pod coordination." @@ -11500,8 +11500,8 @@ class ProxyStartupEvent: Doc: https://docs.datadoghq.com/tracing/trace_collection/automatic_instrumentation/dd_libraries/python/ """ from litellm.litellm_core_utils.dd_tracing import ( - _should_use_dd_profiler, _should_use_dd_tracer, + should_use_dd_profiler, ) if _should_use_dd_tracer(): @@ -11509,7 +11509,7 @@ class ProxyStartupEvent: ddtrace.patch_all(logging=True, openai=False) - if _should_use_dd_profiler(): + if should_use_dd_profiler(): from ddtrace.profiling import Profiler prof: Final = Profiler() @@ -13141,7 +13141,7 @@ async def realtime_websocket_endpoint( await close_after_upstream_handshake_refusal(websocket, e.response.status_code) except Exception as e: verbose_proxy_logger.exception("Internal server error") - redacted_error: Final = _redact_string(str(e)) + redacted_error: Final = redact_string(str(e)) try: await websocket.send_text(realtime_error_event(redacted_error, error_type="server_error")) except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below @@ -14143,7 +14143,7 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) CustomHuggingfaceTokenizer | None, model_info.get("custom_tokenizer", None), ) - _tokenizer_used: Final = await asyncify(litellm.utils._select_tokenizer)( + _tokenizer_used: Final = await asyncify(litellm.utils.select_tokenizer)( model=model_to_use, custom_tokenizer=custom_tokenizer ) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index df27696155d..eae621ba1ce 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -390,11 +390,11 @@ async def responses_api( if data.get("background") and isinstance(response, ResponsesAPIResponse): if response.status in ["queued", "in_progress"]: from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) managed_files_obj: Final = cast( - _PROXY_LiteLLMManagedFiles | None, + PROXY_LiteLLMManagedFiles | None, proxy_logging_obj.get_proxy_hook("managed_files"), ) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 49da000155b..d6d673c6c76 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -10,7 +10,7 @@ from pydantic import TypeAdapter, ValidationError import litellm from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.router_utils.common_utils import _is_proxy_admin_request +from litellm.router_utils.common_utils import is_proxy_admin_request # Client-supplied params that make the router or the call path fabricate a # failure or a delay instead of calling the provider. The ``mock_testing_*`` @@ -64,7 +64,7 @@ def _raise_if_model_fully_blocked(llm_router: LitellmRouter, model_name: object, if not isinstance(llm_router, litellm.Router): return deployments: Final = llm_router.get_model_list(model_name=model_name, team_id=team_id) or [] - if llm_router._are_all_deployments_blocked(deployments): + if llm_router.are_all_deployments_blocked(deployments): raise litellm.PermissionDeniedError( message="Model is blocked", model=model_name, @@ -599,7 +599,7 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr team_id: Final = get_team_id_from_data(data) router_model_names: Final = llm_router.model_names if llm_router is not None else [] - is_proxy_admin_without_team: Final = team_id is None and _is_proxy_admin_request(data) + is_proxy_admin_without_team: Final = team_id is None and is_proxy_admin_request(data) # Preprocess Google GenAI generate content requests if route_type in ["agenerate_content", "agenerate_content_stream"]: @@ -687,7 +687,7 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr if ( deployment and deployment.litellm_params - and not llm_router._is_deployment_blocked(deployment) + and not llm_router.is_deployment_blocked(deployment) ): deployment_creds = deployment.litellm_params.model_dump(exclude_none=True) diff --git a/litellm/proxy/search_endpoints/search_tool_management.py b/litellm/proxy/search_endpoints/search_tool_management.py index 1e6ed40bba3..0b8a5fa46ea 100644 --- a/litellm/proxy/search_endpoints/search_tool_management.py +++ b/litellm/proxy/search_endpoints/search_tool_management.py @@ -200,7 +200,7 @@ async def list_search_tools( } ``` """ - from litellm.litellm_core_utils.litellm_logging import _get_masked_values + from litellm.litellm_core_utils.litellm_logging import get_masked_values from litellm.proxy.proxy_server import prisma_client, proxy_config if prisma_client is None: @@ -229,7 +229,7 @@ async def list_search_tools( tool_name = config_search_tool.get("search_tool_name") if tool_name: litellm_params_dict = dict(config_search_tool.get("litellm_params", {})) - masked_litellm_params_dict = _get_masked_values( + masked_litellm_params_dict = get_masked_values( litellm_params_dict, unmasked_length=4, number_of_asterisks=4, @@ -254,7 +254,7 @@ async def list_search_tools( for db_search_tool in search_tools_from_db: litellm_params_dict = dict(db_search_tool.get("litellm_params", {})) - masked_litellm_params_dict = _get_masked_values( + masked_litellm_params_dict = get_masked_values( litellm_params_dict, unmasked_length=4, number_of_asterisks=4, @@ -529,7 +529,7 @@ async def get_search_tool_info(search_tool_id: str): } ``` """ - from litellm.litellm_core_utils.litellm_logging import _get_masked_values + from litellm.litellm_core_utils.litellm_logging import get_masked_values from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -549,7 +549,7 @@ async def get_search_tool_info(search_tool_id: str): # Mask sensitive data litellm_params_dict: Final = dict(result.get("litellm_params", {})) - masked_litellm_params_dict: Final = _get_masked_values( + masked_litellm_params_dict: Final = get_masked_values( litellm_params_dict, unmasked_length=4, number_of_asterisks=4, diff --git a/litellm/proxy/spend_tracking/savings.py b/litellm/proxy/spend_tracking/savings.py index aa0590ce934..e1e25f26ea8 100644 --- a/litellm/proxy/spend_tracking/savings.py +++ b/litellm/proxy/spend_tracking/savings.py @@ -20,9 +20,9 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.litellm_core_utils.llm_cost_calc.utils import ( - _get_cost_per_unit, calculate_prompt_caching_savings, generic_cost_per_token, + get_cost_per_unit, ) from litellm.types.integrations.anthropic_cache_control_hook import ( GATEWAY_INJECTED_CACHE_METADATA_KEY, @@ -693,7 +693,7 @@ def compute_savings_spend( request_pricing: Final = _request_savings_pricing(model, custom_llm_provider, model_id, llm_router) provider: Final = request_pricing[0] pricing: Final = request_pricing[1] - input_cost: Final = (_get_cost_per_unit(pricing, "input_cost_per_token") or 0.0) if pricing else 0.0 + input_cost: Final = (get_cost_per_unit(pricing, "input_cost_per_token") or 0.0) if pricing else 0.0 compression: Final = max(compression_saved_tokens, 0) * input_cost prompt_caching: Final = _prompt_caching_savings(pricing, provider, usage_object, cost_breakdown, billed_at) or 0.0 gateway_injected_caching: Final = prompt_caching if gateway_injected_cache else 0.0 diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 7c54e154650..af33a3e528b 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -4353,7 +4353,7 @@ async def provider_budgets() -> ProviderBudgetResponse: "No provider budget config found. Please set a provider budget config in the router settings. https://docs.litellm.ai/docs/proxy/provider_budget_routing" ) - router_budget_logger: Final = llm_router._get_router_deployment_budget_limiter() + router_budget_logger: Final = llm_router.get_router_deployment_budget_limiter() if router_budget_logger is None: raise ValueError("No router budget logger found") diff --git a/litellm/proxy/spend_tracking/vantage_endpoints.py b/litellm/proxy/spend_tracking/vantage_endpoints.py index c71105ad283..8ea28f42cd9 100644 --- a/litellm/proxy/spend_tracking/vantage_endpoints.py +++ b/litellm/proxy/spend_tracking/vantage_endpoints.py @@ -410,14 +410,14 @@ async def vantage_dry_run_export( # Use the same pre-transform column names as # FocusExportEngine.dry_run_export_usage_data for consistency. - total_spend: Final = FocusExportEngine._sum_column(data, "spend") - total_tokens: Final = FocusExportEngine._sum_column(data, "total_tokens") + total_spend: Final = FocusExportEngine.sum_column(data, "spend") + total_tokens: Final = FocusExportEngine.sum_column(data, "total_tokens") summary: Final = { "total_records": len(normalized), "total_spend": float(total_spend) if total_spend is not None else 0, "total_tokens": float(total_tokens) if total_tokens is not None else 0, - "unique_teams": FocusExportEngine._count_unique(data, "team_id"), - "unique_models": FocusExportEngine._count_unique(data, "model"), + "unique_teams": FocusExportEngine.count_unique(data, "team_id"), + "unique_models": FocusExportEngine.count_unique(data, "model"), } dry_run_result: Final = { diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index a6e6209298b..b775a1770b3 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -122,7 +122,7 @@ from litellm import ( Router, ) from litellm._internal_context import service_target -from litellm._logging import _redact_string, verbose_proxy_logger +from litellm._logging import redact_string, verbose_proxy_logger from litellm._service_logger import ServiceLogging, ServiceTypes from litellm.caching.caching import DualCache, RedisCache from litellm.caching.dual_cache import LimitedSizeOrderedDict @@ -139,7 +139,7 @@ from litellm.integrations.custom_guardrail import ( from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.prometheus import PrometheusLogger from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting -from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert +from litellm.integrations.SlackAlerting.utils import add_langfuse_trace_id_to_alert from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route from litellm.litellm_core_utils.core_helpers import ( coerce_token_limit, @@ -254,7 +254,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointT from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams from litellm.utils import ( - _add_custom_logger_callback_to_specific_event, # pyright: ignore[reportPrivateUsage] # only string-to-logger helper + add_custom_logger_callback_to_specific_event, ) if TYPE_CHECKING: @@ -350,7 +350,7 @@ def print_verbose(print_statement: object): verbose_proxy_logger.debug("%s\n%s", print_statement, traceback.format_exc()) if litellm.set_verbose: - print(f"LiteLLM Proxy: {_redact_string(str(print_statement))}") # noqa: T201 + print(f"LiteLLM Proxy: {redact_string(str(print_statement))}") # noqa: T201 def _get_email_logger_class(): @@ -1437,9 +1437,9 @@ class ProxyLogging: success_callbacks: Final = tuple(cb for cb in litellm.success_callback if isinstance(cb, str)) failure_callbacks: Final = tuple(cb for cb in litellm.failure_callback if isinstance(cb, str)) for callback in success_callbacks: - _add_custom_logger_callback_to_specific_event(callback, "success") + add_custom_logger_callback_to_specific_event(callback, "success") for callback in failure_callbacks: - _add_custom_logger_callback_to_specific_event(callback, "failure") + add_custom_logger_callback_to_specific_event(callback, "failure") async def update_request_status(self, litellm_call_id: str, status: Literal["success", "fail"]): # only use this if slack alerting is being used @@ -3164,7 +3164,7 @@ class ProxyLogging: extra_kwargs: Final = {} alerting_metadata = {} if request_data is not None: - _url: Final = await _add_langfuse_trace_id_to_alert(request_data=request_data) + _url: Final = await add_langfuse_trace_id_to_alert(request_data=request_data) if _url is not None: extra_kwargs["🪢 Langfuse Trace"] = _url @@ -3211,7 +3211,7 @@ class ProxyLogging: error_message = str(original_exception) if isinstance(traceback_str, str): error_message += traceback_str[:1000] - error_message = _redact_string(error_message) + error_message = redact_string(error_message) asyncio.create_task( self.alerting_handler( message=f"DB read/write call failed: {error_message}", @@ -3285,7 +3285,7 @@ class ProxyLogging: asyncio.create_task( self.alerting_handler( - message=_redact_string(f"LLM API call failed: `{exception_str}`"), + message=redact_string(f"LLM API call failed: `{exception_str}`"), level="High", alert_type=AlertType.llm_exceptions, request_data=request_data, @@ -3315,7 +3315,7 @@ class ProxyLogging: # Remove before callbacks iterate — not serialisable request_data.pop("litellm_logging_obj", None) - redacted_traceback_str: Final = _redact_string(traceback_str) if traceback_str is not None else None + redacted_traceback_str: Final = redact_string(traceback_str) if traceback_str is not None else None # Track the first HTTPException returned or raised by any callback transformed_exception: HTTPException | None = None diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index b67219def38..c2203bf7f67 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -638,5 +638,5 @@ async def index_list( detail=CommonProxyErrors.db_not_connected_error.value, ) - indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db(prisma_client) + indexes: Final = await VectorStoreIndexRegistry.get_vector_store_indexes_from_db(prisma_client) return IndexListResponse(data=indexes) diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index fca591a69ea..4a1f1dc19b8 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -383,7 +383,7 @@ async def list_vector_stores( try: # Get vector stores from database first (source of truth) - vector_stores_from_db: Final = await VectorStoreRegistry._get_vector_stores_from_db(prisma_client=prisma_client) + vector_stores_from_db: Final = await VectorStoreRegistry.get_vector_stores_from_db(prisma_client=prisma_client) # Build map from database vector stores for vector_store in vector_stores_from_db: diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 28814741852..eeee04f82bb 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -357,9 +357,9 @@ async def _arealtime( timeout: float | None = None, query_params: RealtimeQueryParams | None = None, **kwargs, -): +) -> None: """ - Private function to handle the realtime API call. + Handle the realtime API call. For PROXY use only. """ @@ -628,15 +628,15 @@ def _realtime_health_check_auth_headers( return MappingProxyType({"Authorization": f"Bearer {api_key}"}) -async def _realtime_health_check( +async def realtime_health_check( model: str, custom_llm_provider: str, api_key: str | None, api_base: str | None = None, api_version: str | None = None, realtime_protocol: str | None = None, - model_params: dict | None = None, -): + model_params: Mapping[str, object] | None = None, +) -> bool: """ Health check for realtime API - tries connection to the realtime API websocket @@ -740,3 +740,6 @@ async def _realtime_health_check( ssl=ssl_context, ): return True + + +_realtime_health_check = realtime_health_check diff --git a/litellm/responses/litellm_completion_transformation/session_handler.py b/litellm/responses/litellm_completion_transformation/session_handler.py index e6f4b364815..e60a71c7494 100644 --- a/litellm/responses/litellm_completion_transformation/session_handler.py +++ b/litellm/responses/litellm_completion_transformation/session_handler.py @@ -287,7 +287,7 @@ class ResponsesSessionHandler: verbose_proxy_logger.debug("decoding response id=%s", previous_response_id) - decoded_response_id: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(previous_response_id) + decoded_response_id: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(previous_response_id) response_id: Final = decoded_response_id.get("response_id", previous_response_id) if prisma_client is None: return [] diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index e21f706ba60..341324b0c74 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -487,7 +487,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): return buffered def _with_encoded_response_id(self, response: ResponsesAPIResponse) -> ResponsesAPIResponse: - return ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + return ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=response, custom_llm_provider=self.custom_llm_provider, litellm_metadata=self.litellm_metadata, @@ -519,7 +519,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if "text" in self.responses_api_request: response_created_event_data["text"] = self.responses_api_request["text"] response_created_event_data["tool_choice"] = ( - LiteLLMCompletionResponsesConfig._transform_tool_choice_for_responses_api_response( + LiteLLMCompletionResponsesConfig.transform_tool_choice_for_responses_api_response( self.responses_api_request.get("tool_choice") ) ) @@ -760,7 +760,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): annotations: Final = getattr(litellm_complete_object.choices[0].message, "annotations", None) response_annotations: Final = ( - LiteLLMCompletionResponsesConfig._transform_chat_completion_annotations_to_response_output_annotations( + LiteLLMCompletionResponsesConfig.transform_chat_completion_annotations_to_response_output_annotations( annotations=annotations ) ) @@ -787,7 +787,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): annotations = getattr(self.litellm_model_response.choices[0].message, "annotations", None) response_annotations: Final = ( - LiteLLMCompletionResponsesConfig._transform_chat_completion_annotations_to_response_output_annotations( + LiteLLMCompletionResponsesConfig.transform_chat_completion_annotations_to_response_output_annotations( annotations=annotations ) ) @@ -1164,7 +1164,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self.sent_annotation_events = True # Store annotation events to emit them one by one if not hasattr(self, "_pending_annotation_events"): - response_annotations = LiteLLMCompletionResponsesConfig._transform_chat_completion_annotations_to_response_output_annotations( + response_annotations = LiteLLMCompletionResponsesConfig.transform_chat_completion_annotations_to_response_output_annotations( annotations=annotations ) self._pending_annotation_events = [] diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 3a27d21d0ce..cb3c04b6452 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -326,7 +326,7 @@ class LiteLLMCompletionResponsesConfig: return tool_choice @staticmethod - def _transform_tool_choice_for_responses_api_response(tool_choice: object) -> ToolChoice: + def transform_tool_choice_for_responses_api_response(tool_choice: object) -> ToolChoice: if tool_choice is None: return "auto" try: @@ -334,6 +334,8 @@ class LiteLLMCompletionResponsesConfig: except ValidationError: return LiteLLMCompletionResponsesConfig._chat_tool_choice_as_responses_api_tool_choice(tool_choice) + _transform_tool_choice_for_responses_api_response = transform_tool_choice_for_responses_api_response + @staticmethod def _chat_tool_choice_as_responses_api_tool_choice(tool_choice: object) -> ToolChoice: match tool_choice, LiteLLMCompletionResponsesConfig._transform_tool_choice(tool_choice): @@ -2327,7 +2329,7 @@ class LiteLLMCompletionResponsesConfig: return IncompleteDetails(reason=reason) if reason is not None else None @staticmethod - def _tool_call_id_from_responses_item(item_id: str | None, call_id: str | None) -> str: + def tool_call_id_from_responses_item(item_id: str | None, call_id: str | None) -> str: """Bedrock Mantle returns a non-unique, index-based ``call_id`` (``call_0``, ``call_1``, ... that resets every response) alongside a unique ``id`` (``fc_...``). ``call_id`` is the canonical Responses API correlation key, so @@ -2338,6 +2340,8 @@ class LiteLLMCompletionResponsesConfig: return call_id return item_id or call_id or "" + _tool_call_id_from_responses_item = tool_call_id_from_responses_item + @staticmethod def convert_response_function_tool_call_to_chat_completion_tool_call( tool_call_item: object, @@ -2378,7 +2382,7 @@ class LiteLLMCompletionResponsesConfig: function_dict["provider_specific_fields"] = provider_specific_fields tool_call_dict: Final[dict[str, object]] = { - "id": LiteLLMCompletionResponsesConfig._tool_call_id_from_responses_item( + "id": LiteLLMCompletionResponsesConfig.tool_call_id_from_responses_item( getattr(tool_call_item, "id", None), getattr(tool_call_item, "call_id", None), ), @@ -2467,7 +2471,7 @@ class LiteLLMCompletionResponsesConfig: ), parallel_tool_calls=echoed.get("parallel_tool_calls", False), temperature=echoed.get("temperature"), - tool_choice=LiteLLMCompletionResponsesConfig._transform_tool_choice_for_responses_api_response( + tool_choice=LiteLLMCompletionResponsesConfig.transform_tool_choice_for_responses_api_response( responses_api_request.get("tool_choice") ), tools=echoed.get("tools") or [], @@ -2480,7 +2484,7 @@ class LiteLLMCompletionResponsesConfig: ), text=echoed.get("text") or {}, truncation=echoed.get("truncation"), - usage=LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + usage=LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ), user=echoed.get("user"), @@ -2797,7 +2801,7 @@ class LiteLLMCompletionResponsesConfig: ) -> OutputText: annotations: Final = getattr(message, "annotations", None) transformed_annotations: Final = ( - LiteLLMCompletionResponsesConfig._transform_chat_completion_annotations_to_response_output_annotations( + LiteLLMCompletionResponsesConfig.transform_chat_completion_annotations_to_response_output_annotations( annotations=annotations ) ) @@ -2809,7 +2813,7 @@ class LiteLLMCompletionResponsesConfig: ) @staticmethod - def _transform_chat_completion_annotations_to_response_output_annotations( + def transform_chat_completion_annotations_to_response_output_annotations( annotations: list[ChatCompletionAnnotation] | None, ) -> list[GenericResponseOutputItemContentAnnotation]: response_output_annotations: Final[list[GenericResponseOutputItemContentAnnotation]] = [] @@ -2834,8 +2838,12 @@ class LiteLLMCompletionResponsesConfig: return response_output_annotations + _transform_chat_completion_annotations_to_response_output_annotations = ( + transform_chat_completion_annotations_to_response_output_annotations + ) + @staticmethod - def _transform_chat_completion_usage_to_responses_usage( + def transform_chat_completion_usage_to_responses_usage( chat_completion_response: ModelResponse | Usage, ) -> ResponseAPIUsage: if isinstance(chat_completion_response, ModelResponse): @@ -2920,6 +2928,8 @@ class LiteLLMCompletionResponsesConfig: return response_usage + _transform_chat_completion_usage_to_responses_usage = transform_chat_completion_usage_to_responses_usage + @staticmethod def _transform_text_format_to_response_format( text_param: object, diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 5e2793ed40d..f50bdfa42ef 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -202,7 +202,7 @@ async def aresponses_api_with_mcp( ( mcp_tools_with_litellm_proxy, other_tools, - ) = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools) + ) = await LiteLLM_Proxy_MCP_Handler.split_mcp_tools(tools) # Process MCP tools through the complete pipeline (fetch + filter + deduplicate + transform) # Extract user_api_key_auth from litellm_metadata (where it's added by add_user_api_key_auth_to_request_metadata) @@ -220,16 +220,16 @@ async def aresponses_api_with_mcp( ( original_mcp_tools, tool_server_map, - ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + ) = await LiteLLM_Proxy_MCP_Handler.process_mcp_tools_without_openai_transform( user_api_key_auth=user_api_key_auth, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, litellm_trace_id=kwargs.get("litellm_trace_id"), mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, - request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs), + request_tags=LiteLLM_Proxy_MCP_Handler.get_parent_request_tags(kwargs), raw_headers=discovery_raw_headers, ) - openai_tools: Final = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(original_mcp_tools) + openai_tools: Final = LiteLLM_Proxy_MCP_Handler.transform_mcp_tools_to_openai(original_mcp_tools) # Combine with other tools all_tools: Final = openai_tools + other_tools if (openai_tools or other_tools) else None @@ -295,12 +295,12 @@ async def aresponses_api_with_mcp( return mcp_streaming_response # Determine if we should auto-execute tools - should_auto_execute = bool(mcp_tools_with_litellm_proxy) and LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools( + should_auto_execute = bool(mcp_tools_with_litellm_proxy) and LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools( mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy ) # Prepare parameters for the initial call - initial_call_params: Final = LiteLLM_Proxy_MCP_Handler._prepare_initial_call_params( + initial_call_params: Final = LiteLLM_Proxy_MCP_Handler.prepare_initial_call_params( call_params=call_params, should_auto_execute=should_auto_execute ) @@ -322,7 +322,7 @@ async def aresponses_api_with_mcp( # If auto-execute tools is True, then we need to execute the tool calls ######################################################### if should_auto_execute and isinstance(response, ResponsesAPIResponse): - tool_calls: Final = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_response(response=response) + tool_calls: Final = LiteLLM_Proxy_MCP_Handler.extract_tool_calls_from_response(response=response) if tool_calls: user_api_key_auth = kwargs.get("litellm_metadata", {}).get("user_api_key_auth") @@ -338,7 +338,7 @@ async def aresponses_api_with_mcp( tools=tools, ) - tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_results: Final = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map=tool_server_map, served_tools=original_mcp_tools, tool_calls=tool_calls, @@ -349,16 +349,16 @@ async def aresponses_api_with_mcp( raw_headers=raw_headers_from_request, litellm_call_id=kwargs.get("litellm_call_id"), litellm_trace_id=kwargs.get("litellm_trace_id"), - request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs), + request_tags=LiteLLM_Proxy_MCP_Handler.get_parent_request_tags(kwargs), guardrail_context=MCPRequestContext.resolve_guardrail_context( MappingProxyType({**kwargs, "metadata": metadata, "model": model}) ), ) if tool_results: - persistence_disabled: Final = LiteLLM_Proxy_MCP_Handler._is_persistence_disabled(call_params) + persistence_disabled: Final = LiteLLM_Proxy_MCP_Handler.is_persistence_disabled(call_params) - follow_up_input: Final = LiteLLM_Proxy_MCP_Handler._create_follow_up_input( + follow_up_input: Final = LiteLLM_Proxy_MCP_Handler.create_follow_up_input( response=response, tool_results=tool_results, original_input=input, @@ -366,18 +366,18 @@ async def aresponses_api_with_mcp( ) # Prepare parameters for follow-up call (restores original stream setting) - follow_up_call_params: Final = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params( + follow_up_call_params: Final = LiteLLM_Proxy_MCP_Handler.prepare_follow_up_call_params( call_params=call_params, original_stream_setting=stream or False ) # Create tool execution events for streaming if needed tool_execution_events = [] if stream: - tool_execution_events = LiteLLM_Proxy_MCP_Handler._create_tool_execution_events( + tool_execution_events = LiteLLM_Proxy_MCP_Handler.create_tool_execution_events( tool_calls=tool_calls, tool_results=tool_results ) - final_response = await LiteLLM_Proxy_MCP_Handler._make_follow_up_call( + final_response = await LiteLLM_Proxy_MCP_Handler.make_follow_up_call( follow_up_input=follow_up_input, model=model, all_tools=all_tools, @@ -409,15 +409,15 @@ async def aresponses_api_with_mcp( ( mcp_tools_for_output, _, - ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + ) = await LiteLLM_Proxy_MCP_Handler.process_mcp_tools_without_openai_transform( user_api_key_auth=user_api_key_auth, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, - request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs), + request_tags=LiteLLM_Proxy_MCP_Handler.get_parent_request_tags(kwargs), raw_headers=discovery_raw_headers, ) - final_response = LiteLLM_Proxy_MCP_Handler._add_mcp_output_elements_to_response( + final_response = LiteLLM_Proxy_MCP_Handler.add_mcp_output_elements_to_response( response=final_response, mcp_tools_fetched=mcp_tools_for_output, tool_results=tool_results, @@ -741,7 +741,7 @@ async def aresponses( # Update the responses_api_response_id with the model_id if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + response = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=response, litellm_metadata=kwargs.get("litellm_metadata", {}), custom_llm_provider=custom_llm_provider, @@ -1008,7 +1008,7 @@ def _responses_try_dispatch_mcp_gateway( LiteLLM_Proxy_MCP_Handler, ) - if skip_mcp_handler or not LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools): + if skip_mcp_handler or not LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(tools=tools): return None mcp_call_kwargs: Final = { "input": input, @@ -1326,7 +1326,7 @@ def responses( ) reasoning_effort: Final = local_vars.get("reasoning_effort") request_reasoning: Final = ( - LiteLLMResponsesTransformationHandler()._map_reasoning_effort(reasoning_effort) + LiteLLMResponsesTransformationHandler().map_reasoning_effort(reasoning_effort) if current_reasoning is None and reasoning_effort is not None else current_reasoning ) @@ -1418,7 +1418,7 @@ def responses( ) # Decode any litellm-encoded encrypted-content item IDs back to their original IDs - input = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(input) + input = ResponsesAPIRequestUtils.restore_encrypted_content_item_ids_in_input(input) # Call the handler with _is_async flag instead of directly calling the async handler if custom_llm_provider is None: @@ -1446,7 +1446,7 @@ def responses( # Update the responses_api_response_id with the model_id if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + response = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=response, litellm_metadata=kwargs.get("litellm_metadata", {}), custom_llm_provider=custom_llm_provider, @@ -1491,7 +1491,7 @@ async def adelete_responses( kwargs["adelete_responses"] = True # get custom llm provider from response_id - decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils._decode_responses_api_response_id( + decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils.decode_responses_api_response_id( response_id=response_id, ) response_id = decoded_response_id.get("response_id") or response_id @@ -1556,7 +1556,7 @@ def delete_responses( litellm_params: Final = GenericLiteLLMParams(**kwargs) # get custom llm provider from response_id - decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils._decode_responses_api_response_id( + decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils.decode_responses_api_response_id( response_id=response_id, ) response_id = decoded_response_id.get("response_id") or response_id @@ -1649,7 +1649,7 @@ async def aget_responses( kwargs["aget_responses"] = True # get custom llm provider from response_id - decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils._decode_responses_api_response_id( + decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils.decode_responses_api_response_id( response_id=response_id, ) response_id = decoded_response_id.get("response_id") or response_id @@ -1677,7 +1677,7 @@ async def aget_responses( # Update the responses_api_response_id with the model_id if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + response = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=response, litellm_metadata=kwargs.get("litellm_metadata", {}), custom_llm_provider=custom_llm_provider, @@ -1728,7 +1728,7 @@ def get_responses( litellm_params: Final = GenericLiteLLMParams(**kwargs) # get custom llm provider from response_id - decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils._decode_responses_api_response_id( + decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils.decode_responses_api_response_id( response_id=response_id, ) response_id = decoded_response_id.get("response_id") or response_id @@ -1781,7 +1781,7 @@ def get_responses( # Update the responses_api_response_id with the model_id if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + response = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=response, litellm_metadata=kwargs.get("litellm_metadata", {}), custom_llm_provider=custom_llm_provider, @@ -1817,7 +1817,7 @@ async def alist_input_items( loop: Final = asyncio.get_event_loop() kwargs["alist_input_items"] = True - decoded_response_id: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(response_id=response_id) + decoded_response_id: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(response_id=response_id) response_id = decoded_response_id.get("response_id") or response_id custom_llm_provider = decoded_response_id.get("custom_llm_provider") or custom_llm_provider @@ -1876,7 +1876,7 @@ def list_input_items( litellm_params: Final = GenericLiteLLMParams(**kwargs) - decoded_response_id: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(response_id=response_id) + decoded_response_id: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(response_id=response_id) response_id = decoded_response_id.get("response_id") or response_id custom_llm_provider = decoded_response_id.get("custom_llm_provider") or custom_llm_provider @@ -1958,7 +1958,7 @@ async def acancel_responses( kwargs["acancel_responses"] = True # get custom llm provider from response_id - decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils._decode_responses_api_response_id( + decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils.decode_responses_api_response_id( response_id=response_id, ) response_id = decoded_response_id.get("response_id") or response_id @@ -2023,7 +2023,7 @@ def cancel_responses( litellm_params: Final = GenericLiteLLMParams(**kwargs) # get custom llm provider from response_id - decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils._decode_responses_api_response_id( + decoded_response_id: Final[DecodedResponseId] = ResponsesAPIRequestUtils.decode_responses_api_response_id( response_id=response_id, ) response_id = decoded_response_id.get("response_id") or response_id @@ -2146,7 +2146,7 @@ async def acompact_responses( # Update the responses_api_response_id with the model_id if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + response = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=response, litellm_metadata=kwargs.get("litellm_metadata", {}), custom_llm_provider=custom_llm_provider, @@ -2249,7 +2249,7 @@ def compact_responses( # Decode any litellm-encoded encrypted-content item IDs back to their original IDs # before forwarding to the upstream provider. - input = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(input) + input = ResponsesAPIRequestUtils.restore_encrypted_content_item_ids_in_input(input) # Call the handler with _is_async flag instead of directly calling the async handler response = base_llm_http_handler.compact_response_api_handler( @@ -2270,7 +2270,7 @@ def compact_responses( # Update the responses_api_response_id with the model_id if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + response = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=response, litellm_metadata=kwargs.get("litellm_metadata", {}), custom_llm_provider=custom_llm_provider, @@ -2311,7 +2311,7 @@ def _deployment_reasoning_default(kwargs: Mapping[str, object]) -> Reasoning | d return None if isinstance(reasoning_effort, Mapping): return _JSON_OBJECT_ADAPTER.validate_python(reasoning_effort) - return LiteLLMResponsesTransformationHandler()._map_reasoning_effort(reasoning_effort) + return LiteLLMResponsesTransformationHandler().map_reasoning_effort(reasoning_effort) _RESPONSES_WS_ROUTING_HINT_KEYS: Final = frozenset({"input", "previous_response_id"}) diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 8aa4181cc8a..3095a319cf7 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -107,7 +107,7 @@ async def acompletion_with_mcp( ( mcp_tools_with_litellm_proxy, other_tools, - ) = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools) + ) = await LiteLLM_Proxy_MCP_Handler.split_mcp_tools(tools) if not mcp_tools_with_litellm_proxy: # No MCP tools, proceed with regular completion @@ -131,7 +131,7 @@ async def acompletion_with_mcp( ( deduplicated_mcp_tools, tool_server_map, - ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + ) = await LiteLLM_Proxy_MCP_Handler.process_mcp_tools_without_openai_transform( user_api_key_auth=user_api_key_auth, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, litellm_trace_id=context.litellm_trace_id, @@ -141,7 +141,7 @@ async def acompletion_with_mcp( raw_headers=raw_headers, ) - openai_tools: Final = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( + openai_tools: Final = LiteLLM_Proxy_MCP_Handler.transform_mcp_tools_to_openai( deduplicated_mcp_tools, target_format="chat", ) @@ -150,7 +150,7 @@ async def acompletion_with_mcp( all_tools: Final = openai_tools + other_tools if (openai_tools or other_tools) else None # Determine if we should auto-execute tools - should_auto_execute: Final = LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools( + should_auto_execute: Final = LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools( mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy ) @@ -427,13 +427,13 @@ async def acompletion_with_mcp( if isinstance(complete_response, ModelResponse): self.complete_response = complete_response # Extract tool calls from complete response - self.tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response( + self.tool_calls = LiteLLM_Proxy_MCP_Handler.extract_tool_calls_from_chat_response( response=complete_response ) if self.tool_calls: # Execute tool calls - self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + self.tool_results = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map=self.tool_server_map, served_tools=deduplicated_mcp_tools, tool_calls=self.tool_calls, @@ -457,7 +457,7 @@ async def acompletion_with_mcp( return # Create follow-up messages with tool results - follow_up_messages: Final = LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( + follow_up_messages: Final = LiteLLM_Proxy_MCP_Handler.create_follow_up_messages_for_chat( original_messages=self.messages, response=self.complete_response, tool_results=self.tool_results, @@ -597,7 +597,7 @@ async def acompletion_with_mcp( return initial_response # Extract tool calls from response - tool_calls: Final = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(response=initial_response) + tool_calls: Final = LiteLLM_Proxy_MCP_Handler.extract_tool_calls_from_chat_response(response=initial_response) if not tool_calls: _add_mcp_metadata_to_response( @@ -607,7 +607,7 @@ async def acompletion_with_mcp( return initial_response # Execute tool calls - tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_results: Final = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map=tool_server_map, tool_calls=tool_calls, served_tools=deduplicated_mcp_tools, @@ -631,7 +631,7 @@ async def acompletion_with_mcp( return initial_response # Create follow-up messages with tool results - follow_up_messages: Final = LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( + follow_up_messages: Final = LiteLLM_Proxy_MCP_Handler.create_follow_up_messages_for_chat( original_messages=messages, response=initial_response, tool_results=tool_results, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 1d6c0695621..ba8e236b405 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -136,7 +136,7 @@ class LiteLLM_Proxy_MCP_Handler: """ @staticmethod - def _get_parent_request_tags(kwargs: dict[str, Any] | None) -> list[str]: + def get_parent_request_tags(kwargs: dict[str, Any] | None) -> list[str]: """Tags from the parent LLM request, using the same extraction logic as standard logging (incl. User-Agent).""" if not kwargs: return [] @@ -144,17 +144,21 @@ class LiteLLM_Proxy_MCP_Handler: litellm_params: Final = kwargs.get("litellm_params") or kwargs proxy_server_request = litellm_params.get("proxy_server_request") or kwargs.get("proxy_server_request") or {} - return StandardLoggingPayloadSetup._get_request_tags( + return StandardLoggingPayloadSetup.get_request_tags( litellm_params=litellm_params, proxy_server_request=proxy_server_request, ) + _get_parent_request_tags = get_parent_request_tags + @staticmethod - def _should_use_litellm_mcp_gateway(tools: Iterable[ToolParam] | None) -> bool: + def should_use_litellm_mcp_gateway(tools: Iterable[ToolParam] | None) -> bool: """True when a tool may name this gateway: server_url "litellm_proxy..." or an http(s) URL ending in /mcp/. `_split_mcp_tools` then settles which of the latter the gateway actually serves.""" return any(_names_gateway_explicitly(tool) or _proxy_path_mcp_name(tool) is not None for tool in tools or ()) + _should_use_litellm_mcp_gateway = should_use_litellm_mcp_gateway + @staticmethod def _parse_mcp_tools(tools: Iterable[Mapping[str, object]] | None) -> SplitTools: items: Final = tuple(tools or ()) @@ -163,7 +167,7 @@ class LiteLLM_Proxy_MCP_Handler: return gateway_tools, other_tools @staticmethod - async def _split_mcp_tools( + async def split_mcp_tools( tools: Iterable[Mapping[str, object]] | None, served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] = _gateway_served_names, ) -> SplitTools: @@ -178,6 +182,8 @@ class LiteLLM_Proxy_MCP_Handler: ] ) + _split_mcp_tools = split_mcp_tools + @staticmethod async def routes_through_gateway( tools: Iterable[Mapping[str, object]] | None, @@ -414,7 +420,7 @@ class LiteLLM_Proxy_MCP_Handler: return filtered_tools @staticmethod - async def _process_mcp_tools_to_openai_format( + async def process_mcp_tools_to_openai_format( user_api_key_auth: "UserAPIKeyAuth | None", mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]], litellm_trace_id: str | None = None, @@ -435,19 +441,21 @@ class LiteLLM_Proxy_MCP_Handler: ( deduplicated_mcp_tools, tool_server_map, - ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + ) = await LiteLLM_Proxy_MCP_Handler.process_mcp_tools_without_openai_transform( user_api_key_auth, mcp_tools_with_litellm_proxy, litellm_trace_id=litellm_trace_id, request_tags=request_tags, ) - openai_tools: Final = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(deduplicated_mcp_tools) + openai_tools: Final = LiteLLM_Proxy_MCP_Handler.transform_mcp_tools_to_openai(deduplicated_mcp_tools) return openai_tools, tool_server_map + _process_mcp_tools_to_openai_format = process_mcp_tools_to_openai_format + @staticmethod - async def _process_mcp_tools_without_openai_transform( + async def process_mcp_tools_without_openai_transform( user_api_key_auth: "UserAPIKeyAuth | None", mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]], litellm_trace_id: str | None = None, @@ -502,22 +510,24 @@ class LiteLLM_Proxy_MCP_Handler: return deduplicated_mcp_tools, tool_server_map + _process_mcp_tools_without_openai_transform = process_mcp_tools_without_openai_transform + @overload @staticmethod - def _transform_mcp_tools_to_openai( + def transform_mcp_tools_to_openai( mcp_tools: Sequence[MCPTool], target_format: Literal["responses"] = ..., ) -> list[FunctionToolParam]: ... @overload @staticmethod - def _transform_mcp_tools_to_openai( + def transform_mcp_tools_to_openai( mcp_tools: Sequence[MCPTool], target_format: Literal["chat"], ) -> list[ChatCompletionToolParam]: ... @staticmethod - def _transform_mcp_tools_to_openai( + def transform_mcp_tools_to_openai( mcp_tools: Sequence[MCPTool], target_format: Literal["responses", "chat"] = "responses", ) -> Sequence[FunctionToolParam | ChatCompletionToolParam]: @@ -536,8 +546,10 @@ class LiteLLM_Proxy_MCP_Handler: return openai_tools + _transform_mcp_tools_to_openai = transform_mcp_tools_to_openai + @staticmethod - def _should_auto_execute_tools( + def should_auto_execute_tools( mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]], ) -> bool: """Check if we should auto-execute tool calls. @@ -562,8 +574,10 @@ class LiteLLM_Proxy_MCP_Handler: return False return True + _should_auto_execute_tools = should_auto_execute_tools + @staticmethod - def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> list[object]: + def extract_tool_calls_from_response(response: ResponsesAPIResponse) -> list[object]: """Extract tool calls from the response output.""" tool_calls: Final[list[object]] = [] for output_item in response.output: @@ -576,8 +590,10 @@ class LiteLLM_Proxy_MCP_Handler: return tool_calls + _extract_tool_calls_from_response = extract_tool_calls_from_response + @staticmethod - def _extract_tool_calls_from_chat_response(response: ModelResponse) -> list[object]: + def extract_tool_calls_from_chat_response(response: ModelResponse) -> list[object]: """Extract tool calls from a chat completion response.""" tool_calls: Final[list[object]] = [] @@ -600,8 +616,10 @@ class LiteLLM_Proxy_MCP_Handler: return tool_calls + _extract_tool_calls_from_chat_response = extract_tool_calls_from_chat_response + @staticmethod - def _extract_tool_call_details( + def extract_tool_call_details( tool_call: object, ) -> tuple[str | None, str | None, str | None]: """Extract tool name, arguments, and call_id from a tool call.""" @@ -634,6 +652,8 @@ class LiteLLM_Proxy_MCP_Handler: return tool_name, tool_arguments, tool_call_id + _extract_tool_call_details = extract_tool_call_details + @staticmethod def _parse_tool_arguments(tool_arguments: str | None) -> dict[str, object]: """Parse tool arguments, handling both string and dict formats.""" @@ -684,7 +704,7 @@ class LiteLLM_Proxy_MCP_Handler: return result_text or "Tool executed successfully" @staticmethod - async def _execute_tool_calls( + async def execute_tool_calls( tool_server_map: dict[str, str], tool_calls: Sequence[object], user_api_key_auth: "UserAPIKeyAuth | None", @@ -724,7 +744,7 @@ class LiteLLM_Proxy_MCP_Handler: tool_name, tool_arguments, tool_call_id, - ) = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) + ) = LiteLLM_Proxy_MCP_Handler.extract_tool_call_details(tool_call) if not tool_name: verbose_logger.warning("Tool call missing name: %s", tool_call) @@ -977,8 +997,10 @@ class LiteLLM_Proxy_MCP_Handler: return tool_results + _execute_tool_calls = execute_tool_calls + @staticmethod - def _create_follow_up_messages_for_chat( + def create_follow_up_messages_for_chat( original_messages: list[object], response: ModelResponse, tool_results: Sequence[Mapping[str, object]], @@ -1022,13 +1044,17 @@ class LiteLLM_Proxy_MCP_Handler: return follow_up_messages + _create_follow_up_messages_for_chat = create_follow_up_messages_for_chat + @staticmethod - def _is_persistence_disabled(call_params: Mapping[str, object]) -> bool: + def is_persistence_disabled(call_params: Mapping[str, object]) -> bool: """store=false means the provider kept nothing, so the follow-up call cannot chain on a response id.""" return call_params.get("store") is False + _is_persistence_disabled = is_persistence_disabled + @staticmethod - def _create_follow_up_input( + def create_follow_up_input( response: ResponsesAPIResponse, tool_results: Sequence[Mapping[str, object]], original_input: str | ResponseInputParam | None = None, @@ -1106,8 +1132,10 @@ class LiteLLM_Proxy_MCP_Handler: return follow_up_input + _create_follow_up_input = create_follow_up_input + @staticmethod - async def _make_follow_up_call( + async def make_follow_up_call( follow_up_input: list[Any], model: str, all_tools: Sequence[ResponsesToolParam] | None, @@ -1123,6 +1151,8 @@ class LiteLLM_Proxy_MCP_Handler: **call_params, ) + _make_follow_up_call = make_follow_up_call + @staticmethod async def _log_mcp_tool_failure( *, @@ -1230,7 +1260,7 @@ class LiteLLM_Proxy_MCP_Handler: return request_params @staticmethod - def _create_tool_execution_events( + def create_tool_execution_events( tool_calls: Sequence[object], tool_results: Sequence[MCPToolResult] ) -> list[ResponsesAPIStreamingResponse]: """ @@ -1261,7 +1291,7 @@ class LiteLLM_Proxy_MCP_Handler: name, args, call_id, - ) = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) + ) = LiteLLM_Proxy_MCP_Handler.extract_tool_call_details(tool_call) if call_id == tool_call_id: tool_name = name or "unknown" tool_arguments = args or "{}" @@ -1279,8 +1309,10 @@ class LiteLLM_Proxy_MCP_Handler: return tool_execution_events + _create_tool_execution_events = create_tool_execution_events + @staticmethod - def _prepare_initial_call_params(call_params: Mapping[str, object], should_auto_execute: bool) -> dict[str, Any]: + def prepare_initial_call_params(call_params: Mapping[str, object], should_auto_execute: bool) -> dict[str, Any]: """ Prepare call parameters for the initial LLM call. @@ -1295,8 +1327,10 @@ class LiteLLM_Proxy_MCP_Handler: return initial_params + _prepare_initial_call_params = prepare_initial_call_params + @staticmethod - def _prepare_follow_up_call_params( + def prepare_follow_up_call_params( call_params: Mapping[str, object], original_stream_setting: bool ) -> dict[str, Any]: """ @@ -1315,8 +1349,10 @@ class LiteLLM_Proxy_MCP_Handler: return follow_up_params + _prepare_follow_up_call_params = prepare_follow_up_call_params + @staticmethod - def _add_mcp_output_elements_to_response( + def add_mcp_output_elements_to_response( response: ResponsesAPIResponse, mcp_tools_fetched: Sequence[object], tool_results: Sequence[Mapping[str, object]], @@ -1363,3 +1399,5 @@ class LiteLLM_Proxy_MCP_Handler: response.output.append(tool_results_output.model_dump()) return response + + _add_mcp_output_elements_to_response = add_mcp_output_elements_to_response diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index b1f12233f33..d486db72304 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -416,7 +416,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): LiteLLM_Proxy_MCP_Handler, ) - return LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(self.mcp_tools_with_litellm_proxy) + return LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools(self.mcp_tools_with_litellm_proxy) def _make_stream_error_event(self) -> ResponsesAPIStreamingResponse: err: Final = self._stream_error @@ -741,7 +741,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): try: # Extract tool calls from the response if self.collected_response is not None: - tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_response(self.collected_response) + tool_calls = LiteLLM_Proxy_MCP_Handler.extract_tool_calls_from_response(self.collected_response) else: tool_calls = [] if not tool_calls: @@ -759,7 +759,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): tool_name, tool_arguments, tool_call_id, - ) = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) + ) = LiteLLM_Proxy_MCP_Handler.extract_tool_call_details(tool_call) if tool_name and tool_call_id: item_id = f"mcp_{uuid.uuid4().hex[:8]}" output_index = next_output_index @@ -796,7 +796,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.tool_execution_events.extend(call_events[:-1]) # Execute the tools - tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_results: Final = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map=self.tool_server_map, served_tools=self.served_tools, tool_calls=tool_calls, @@ -807,7 +807,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): raw_headers=self.raw_headers, litellm_call_id=self.litellm_call_id, litellm_trace_id=self.litellm_trace_id, - request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(self.original_request_params), + request_tags=LiteLLM_Proxy_MCP_Handler.get_parent_request_tags(self.original_request_params), guardrail_context=MCPRequestContext.resolve_guardrail_context(self.original_request_params), ) @@ -824,7 +824,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): name, args, call_id, - ) = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) + ) = LiteLLM_Proxy_MCP_Handler.extract_tool_call_details(tool_call) if call_id == tool_call_id: tool_name = name or "unknown" tool_arguments = args or "{}" @@ -900,11 +900,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): try: # Create follow-up input if self.collected_response is not None: - persistence_disabled: Final = LiteLLM_Proxy_MCP_Handler._is_persistence_disabled( + persistence_disabled: Final = LiteLLM_Proxy_MCP_Handler.is_persistence_disabled( self.original_request_params ) - follow_up_input: Final = LiteLLM_Proxy_MCP_Handler._create_follow_up_input( + follow_up_input: Final = LiteLLM_Proxy_MCP_Handler.create_follow_up_input( response=self.collected_response, tool_results=self.tool_results, original_input=self.original_request_params.get("input"), diff --git a/litellm/responses/mcp/request_context.py b/litellm/responses/mcp/request_context.py index 0bbed24b6ac..6628219e280 100644 --- a/litellm/responses/mcp/request_context.py +++ b/litellm/responses/mcp/request_context.py @@ -8,7 +8,7 @@ surface from silently dropping a field: omitting the auth headers, for instance, still executes the tool, just with no credentials. """ -from collections.abc import Iterable, Mapping, Sequence +from collections.abc import Iterable, Mapping from copy import deepcopy from dataclasses import dataclass from types import MappingProxyType @@ -33,10 +33,10 @@ class MCPRequestContext: user_api_key_auth: "UserAPIKeyAuth | None" mcp_auth_header: str | None = None - mcp_server_auth_headers: Mapping[str, Mapping[str, str]] | None = None - oauth2_headers: Mapping[str, str] | None = None - raw_headers: Mapping[str, str] | None = None - request_tags: Sequence[str] | None = None + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None + oauth2_headers: dict[str, str] | None = None + raw_headers: dict[str, str] | None = None + request_tags: list[str] | None = None litellm_trace_id: str | None = None litellm_call_id: str | None = None guardrail_context: Mapping[str, object] | None = None @@ -83,7 +83,7 @@ class MCPRequestContext: mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, - request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(dict(kwargs)), + request_tags=LiteLLM_Proxy_MCP_Handler.get_parent_request_tags(dict(kwargs)), litellm_trace_id=kwargs.get("litellm_trace_id"), litellm_call_id=kwargs.get("litellm_call_id"), guardrail_context=cls.resolve_guardrail_context(kwargs), diff --git a/litellm/responses/sse_output_recovery.py b/litellm/responses/sse_output_recovery.py index adc6a30319c..cdabcb504d3 100644 --- a/litellm/responses/sse_output_recovery.py +++ b/litellm/responses/sse_output_recovery.py @@ -29,7 +29,7 @@ def parse_sse_json_chunk(chunk: str) -> dict[str, object] | None: # Import locally to avoid a circular import with the streaming handler. from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - stripped_chunk: Final = (CustomStreamWrapper._strip_sse_data_from_chunk(chunk.strip()) or "").strip() + stripped_chunk: Final = (CustomStreamWrapper.strip_sse_data_from_chunk(chunk.strip()) or "").strip() if not stripped_chunk or stripped_chunk == STREAM_SSE_DONE_STRING or stripped_chunk.startswith("event:"): return None try: diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index eb7a6556386..fd88a4a5de4 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -6,11 +6,11 @@ import json import time import traceback import uuid -from collections.abc import Awaitable, Callable, Coroutine, Iterable, Mapping, Sequence +from collections.abc import AsyncIterable, Awaitable, Callable, Coroutine, Iterable, Mapping, Sequence from datetime import datetime from functools import lru_cache from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Protocol, overload, runtime_checkable +from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Protocol, cast, overload, runtime_checkable import httpx from openai._streaming import SSEDecoder @@ -71,12 +71,15 @@ class ProjectQuotaCallback(Protocol): @lru_cache(maxsize=1) -def _get_openai_response_types(): +def get_openai_response_types(): from litellm.types.llms import openai as openai_types return openai_types +_get_openai_response_types = get_openai_response_types + + def _is_json_object(value: object) -> TypeIs[dict[str, object]]: # guard-ok: trivial isinstance; JSON keys are str return isinstance(value, dict) @@ -376,7 +379,7 @@ class BaseResponsesAPIStreamingIterator: return None if self.logging_obj.completion_start_time is None: - self.logging_obj._update_completion_start_time(completion_start_time=datetime.now()) + self.logging_obj.update_completion_start_time(completion_start_time=datetime.now()) try: # Parse the JSON chunk @@ -400,7 +403,7 @@ class BaseResponsesAPIStreamingIterator: openai_responses_api_chunk, "response", None ) if response_object is not None: - response: Final = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + response: Final = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=response_object, litellm_metadata=self.litellm_metadata, custom_llm_provider=self.custom_llm_provider, @@ -424,7 +427,7 @@ class BaseResponsesAPIStreamingIterator: ): _item: Final[object] = getattr(openai_responses_api_chunk, "item", None) if _item is not None: - ResponsesAPIRequestUtils._encode_container_id_on_output_item( + ResponsesAPIRequestUtils.encode_container_id_on_output_item( item=_item, custom_llm_provider=self.custom_llm_provider, model_id=_stream_model_id, @@ -432,7 +435,7 @@ class BaseResponsesAPIStreamingIterator: elif _event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED: _annotation: Final[object] = getattr(openai_responses_api_chunk, "annotation", None) if _annotation is not None: - ResponsesAPIRequestUtils._encode_container_id_on_output_item( + ResponsesAPIRequestUtils.encode_container_id_on_output_item( item=_annotation, custom_llm_provider=self.custom_llm_provider, model_id=_stream_model_id, @@ -443,13 +446,13 @@ class BaseResponsesAPIStreamingIterator: ) if _part is not None: if isinstance(_part, dict): - ResponsesAPIRequestUtils._encode_container_ids_in_annotations( + ResponsesAPIRequestUtils.encode_container_ids_in_annotations( _part.get("annotations"), self.custom_llm_provider, _stream_model_id, ) else: - ResponsesAPIRequestUtils._encode_container_ids_in_annotations( + ResponsesAPIRequestUtils.encode_container_ids_in_annotations( getattr(_part, "annotations", None), self.custom_llm_provider, _stream_model_id, @@ -457,7 +460,7 @@ class BaseResponsesAPIStreamingIterator: # Wrap encrypted_content in streaming events (output_item.added, output_item.done) if self.litellm_metadata and self.litellm_metadata.get("encrypted_content_affinity_enabled"): - openai_types = _get_openai_response_types() + openai_types = get_openai_response_types() event_type: Final = getattr(openai_responses_api_chunk, "type", None) if event_type in ( openai_types.ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, @@ -470,7 +473,7 @@ class BaseResponsesAPIStreamingIterator: model_id: Final = _model_id_from_metadata(self.litellm_metadata) if model_id: wrapped_content: Final = ( - ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( + ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id( encrypted_content, model_id ) ) @@ -478,7 +481,7 @@ class BaseResponsesAPIStreamingIterator: # Store the completed response (also for incomplete/failed so logging still fires) _chunk_type: Final = getattr(openai_responses_api_chunk, "type", None) - openai_types = _get_openai_response_types() + openai_types = get_openai_response_types() if openai_responses_api_chunk and _chunk_type in ( openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, @@ -632,7 +635,7 @@ class BaseResponsesAPIStreamingIterator: return try: self.logging_obj.model_call_details["combined_usage_object"] = ( - ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage_obj) + ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage_obj) ) except (TypeError, ValueError) as usage_error: verbose_logger.debug( @@ -641,7 +644,7 @@ class BaseResponsesAPIStreamingIterator: ) return self.logging_obj.model_call_details["response_cost"] = ( - self.logging_obj._response_cost_calculator(result=response_obj) or 0.0 + self.logging_obj.response_cost_calculator(result=response_obj) or 0.0 ) def _map_error_event_exception(self, error_obj: object) -> Exception: @@ -671,7 +674,7 @@ class BaseResponsesAPIStreamingIterator: ) def _get_completed_response_object(self) -> ResponsesAPIResponse | None: - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() completed_response: Final = self.completed_response if isinstance(completed_response, openai_types.ResponsesAPIResponse): return completed_response @@ -687,7 +690,7 @@ class BaseResponsesAPIStreamingIterator: return completed_response: Final = self.completed_response - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() if getattr(completed_response, "type", None) != openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: return @@ -697,7 +700,7 @@ class BaseResponsesAPIStreamingIterator: if response_obj is None or is_response_without_output(response_obj): return - caching_handler: Final[LLMCachingHandler | None] = getattr(self.logging_obj, "_llm_caching_handler", None) + caching_handler: Final[LLMCachingHandler | None] = getattr(self.logging_obj, "llm_caching_handler", None) if caching_handler is None: return @@ -1238,7 +1241,7 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise StopAsyncIteration evt: Final = self._events[self._idx] self._idx += 1 - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: self.completed_response = evt self._log_completed_response(is_async=True) @@ -1252,7 +1255,7 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise StopIteration evt: Final = self._events[self._idx] self._idx += 1 - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: self.completed_response = evt self._log_completed_response(is_async=False) @@ -1305,7 +1308,7 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise StopAsyncIteration evt: Final = self._events[self._idx] self._idx += 1 - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: self.completed_response = evt self._log_completed_response(is_async=True) @@ -1319,7 +1322,7 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise StopIteration evt: Final = self._events[self._idx] self._idx += 1 - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: self.completed_response = evt self._log_completed_response(is_async=False) @@ -1351,7 +1354,7 @@ def _build_response_status_event( ], transformed: ResponsesAPIResponse, ) -> ResponsesAPIStreamingResponse: - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() in_progress_response: Final = transformed.model_copy( deep=True, update={"status": "in_progress", "output": []}, @@ -1368,7 +1371,7 @@ def _build_content_part_done_event( content_index: int, part_payload: Mapping[str, object], ) -> ResponsesAPIStreamingResponse | None: - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() part_type: Final = part_payload.get("type") part: PART_UNION_TYPES if part_type == "output_text": @@ -1412,7 +1415,7 @@ def _add_text_like_part_events( part_payload: Mapping[str, object], chunk_size: int, ) -> None: - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() part_type: Final = part_payload.get("type") if part_type == "output_text": text: Final = str(part_payload.get("text") or "") @@ -1591,7 +1594,7 @@ def _stamp_responses_usage_cost( if isinstance(getattr(usage_obj, "cost", None), (int, float)): return try: - cost: Final[float | None] = logging_obj._response_cost_calculator(result=response_obj) + cost: Final[float | None] = logging_obj.response_cost_calculator(result=response_obj) except Exception: return if isinstance(cost, (int, float)) and cost > 0: @@ -1604,7 +1607,7 @@ def build_synthetic_response_events( logging_obj: LiteLLMLoggingObj | None, chunk_size: int, ) -> list[ResponsesAPIStreamingResponse]: - openai_types: Final = _get_openai_response_types() + openai_types: Final = get_openai_response_types() _stamp_responses_usage_cost(transformed, logging_obj) events: Final[list[ResponsesAPIStreamingResponse]] = [ @@ -1839,7 +1842,7 @@ def _ws_event_error(event: Mapping[str, object]) -> object: def _restore_input_item_ids(items: Sequence[object]) -> Sequence[object]: - return ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(copy.deepcopy(list(items))) # pyright: ignore[reportPrivateUsage] # same restore the HTTP responses path runs + return ResponsesAPIRequestUtils.restore_encrypted_content_item_ids_in_input(copy.deepcopy(list(items))) def _restored_container_fields(container: Mapping[str, object]) -> Mapping[str, object]: @@ -1880,7 +1883,7 @@ def _wrap_output_item_encrypted_content( encrypted_content: Final = item.get("encrypted_content") if not isinstance(encrypted_content, str) or not encrypted_content: return None - wrapped_content: Final = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( # pyright: ignore[reportPrivateUsage] # same wrap the HTTP streaming path applies + wrapped_content: Final = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id( encrypted_content=encrypted_content, model_id=model_id ) return {**event_obj, "item": {**item, "encrypted_content": wrapped_content}} @@ -2029,7 +2032,7 @@ class ResponsesWebSocketStreaming: logging_result: Final = LiteLLMRealtimeStreamLoggingObject( usage=usage, results=self.messages, service_tier=service_tier ) - response_cost: Final = self.logging_obj._response_cost_calculator(result=logging_result) or 0.0 # pyright: ignore[reportPrivateUsage] # as the HTTP streaming iterator does + response_cost: Final = self.logging_obj.response_cost_calculator(result=logging_result) or 0.0 self.logging_obj.record_partial_usage_for_failure(usage, response_cost) def _wrap_response_event(self, response_str: str) -> str: @@ -2039,7 +2042,7 @@ class ResponsesWebSocketStreaming: return response_str response: Final = event_obj.get("response") if _is_json_object(response): - wrapped_response: Final = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( # pyright: ignore[reportPrivateUsage] # same wrap the HTTP streaming path applies + wrapped_response: Final = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=response, custom_llm_provider=self.custom_llm_provider, litellm_metadata=self.litellm_metadata, @@ -2484,8 +2487,8 @@ class ResponsesWebSocketStreaming: # --------------------------------------------------------------------------- _RESPONSE_CREATE_PARAMS: Final[frozenset[str]] = ( - _get_openai_response_types().ResponsesAPIRequestParams.__required_keys__ - | _get_openai_response_types().ResponsesAPIRequestParams.__optional_keys__ + get_openai_response_types().ResponsesAPIRequestParams.__required_keys__ + | get_openai_response_types().ResponsesAPIRequestParams.__optional_keys__ ) _MANAGED_WS_SKIP_KWARGS: Final[frozenset[str]] = frozenset( @@ -2591,7 +2594,7 @@ class ManagedResponsesWebSocketHandler: The key is the *decoded* response ID (the raw provider response ID before LiteLLM base64-encodes it into the ``resp_...`` format). """ - decoded: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(previous_response_id) + decoded: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(previous_response_id) raw_id: Final = decoded.get("response_id", previous_response_id) return list(self._session_history.get(raw_id, [])) @@ -2615,7 +2618,7 @@ class ManagedResponsesWebSocketHandler: encoded_id: Final[str | None] = raw_id if isinstance(raw_id, str) else None if not encoded_id: return None - decoded: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(encoded_id) + decoded: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(encoded_id) return decoded.get("response_id", encoded_id) @staticmethod @@ -2701,7 +2704,7 @@ class ManagedResponsesWebSocketHandler: """Return True for synthetic warmup IDs that only exist on this connection.""" if not response_id: return False - decoded: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(response_id) + decoded: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(response_id) raw_id: Final = decoded.get("response_id", response_id) return str(raw_id).startswith(_WARMUP_RESPONSE_ID_PREFIX) @@ -2871,7 +2874,7 @@ class ManagedResponsesWebSocketHandler: """ terminal_event: _MutableJsonObject | None = None stream_response: Final = await litellm.aresponses(model=model, **call_kwargs) - async for chunk in stream_response: + async for chunk in cast(AsyncIterable[object], stream_response): # cast-ok: aresponses returns an async stream if chunk is None: continue # Read type from the object before serializing to avoid double JSON parse diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index a97c4392c09..3e308b73d12 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -4,7 +4,7 @@ from collections.abc import Callable, Iterable, Mapping, Sequence from functools import reduce from typing import Any, Final, Optional, TypeVar, Union, cast, get_type_hints, overload -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter from typing_extensions import TypeIs # noqa: TID251 # narrows untyped wire payloads without a runtime conversion import litellm @@ -214,7 +214,7 @@ class ResponsesAPIRequestUtils: Returns: A dictionary of supported parameters for the responses API """ - from litellm.utils import _apply_openai_param_overrides + from litellm.utils import apply_openai_param_overrides # Remove None values and internal parameters # Get supported parameters for the model @@ -233,20 +233,23 @@ class ResponsesAPIRequestUtils: ) # Map parameters to provider-specific format - mapped_params: Final = responses_api_provider_config.map_openai_params( - response_api_optional_params=response_api_optional_params, - model=model, - drop_params=should_drop_params, + mapped_params: Final[dict[str, object]] = TypeAdapter(dict[str, object]).validate_python( + responses_api_provider_config.map_openai_params( + response_api_optional_params=response_api_optional_params, + model=model, + drop_params=should_drop_params, + ), + strict=True, ) stream_options: Final = normalize_responses_api_stream_options(mapped_params.get("stream_options")) - params_with_normalized_stream_options: Final = { + params_with_normalized_stream_options: Final[dict[str, object]] = { **{key: value for key, value in mapped_params.items() if key != "stream_options"}, **({} if stream_options is None else {"stream_options": stream_options}), } # add any allowed_openai_params to the mapped_params - return _apply_openai_param_overrides( + return apply_openai_param_overrides( optional_params=params_with_normalized_stream_options, non_default_params=non_default_params, allowed_openai_params=allowed_openai_params or [], @@ -309,7 +312,7 @@ class ResponsesAPIRequestUtils: # fmt: off @overload @staticmethod - def _update_responses_api_response_id_with_model_id( + def update_responses_api_response_id_with_model_id( responses_api_response: ResponsesAPIResponse, custom_llm_provider: str | None, litellm_metadata: dict[str, object] | None = None, @@ -318,7 +321,7 @@ class ResponsesAPIRequestUtils: @overload @staticmethod - def _update_responses_api_response_id_with_model_id( + def update_responses_api_response_id_with_model_id( responses_api_response: dict[str, object], custom_llm_provider: str | None, litellm_metadata: dict[str, object] | None = None, @@ -328,7 +331,7 @@ class ResponsesAPIRequestUtils: # fmt: on @staticmethod - def _update_responses_api_response_id_with_model_id( + def update_responses_api_response_id_with_model_id( responses_api_response: ResponsesAPIResponse | dict[str, Any], custom_llm_provider: str | None, litellm_metadata: dict[str, Any] | None = None, @@ -381,6 +384,8 @@ class ResponsesAPIRequestUtils: return responses_api_response + _update_responses_api_response_id_with_model_id = update_responses_api_response_id_with_model_id + @staticmethod def _build_encrypted_item_id(model_id: str, item_id: str) -> str: """Encode model_id into an output item ID for encrypted-content items. @@ -392,7 +397,7 @@ class ResponsesAPIRequestUtils: return f"encitem_{encoded}" @staticmethod - def _decode_encrypted_item_id(encoded_id: str) -> dict[str, str] | None: + def decode_encrypted_item_id(encoded_id: str) -> dict[str, str] | None: """Decode a litellm-encoded encrypted-content item ID. Returns a dict with ``model_id`` and ``item_id`` keys, or ``None`` if @@ -417,8 +422,10 @@ class ResponsesAPIRequestUtils: except Exception: return None + _decode_encrypted_item_id = decode_encrypted_item_id + @staticmethod - def _wrap_encrypted_content_with_model_id(encrypted_content: str, model_id: str) -> str: + def wrap_encrypted_content_with_model_id(encrypted_content: str, model_id: str) -> str: """Wrap encrypted_content with model_id metadata for affinity routing. When Codex or other clients send items with encrypted_content but no ID, @@ -430,8 +437,10 @@ class ResponsesAPIRequestUtils: encoded_metadata: Final = base64.b64encode(metadata.encode("utf-8")).decode("utf-8") return f"litellm_enc:{encoded_metadata};{encrypted_content}" + _wrap_encrypted_content_with_model_id = wrap_encrypted_content_with_model_id + @staticmethod - def _unwrap_encrypted_content_with_model_id( + def unwrap_encrypted_content_with_model_id( wrapped_content: str, ) -> tuple[str | None, str]: """Unwrap encrypted_content to extract model_id and original content. @@ -463,6 +472,8 @@ class ResponsesAPIRequestUtils: except Exception: return None, wrapped_content + _unwrap_encrypted_content_with_model_id = unwrap_encrypted_content_with_model_id + @staticmethod def _update_encrypted_content_item_ids_in_response( response: Union["ResponsesAPIResponse", dict[str, object]], @@ -495,7 +506,7 @@ class ResponsesAPIRequestUtils: if encrypted_content and isinstance(encrypted_content, str): # Always wrap encrypted_content with model_id for redundancy - item["encrypted_content"] = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( + item["encrypted_content"] = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id( encrypted_content, model_id ) # Also encode the ID if present @@ -508,7 +519,7 @@ class ResponsesAPIRequestUtils: if encrypted_content and isinstance(encrypted_content, str): # Always wrap encrypted_content with model_id for redundancy try: - item.encrypted_content = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( + item.encrypted_content = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id( encrypted_content, model_id ) except AttributeError: @@ -523,7 +534,7 @@ class ResponsesAPIRequestUtils: return response @staticmethod - def _restore_encrypted_content_item_ids_in_input(request_input: _RequestInputT) -> _RequestInputT: + def restore_encrypted_content_item_ids_in_input(request_input: _RequestInputT) -> _RequestInputT: """Decode litellm-encoded item IDs in request input back to original IDs. Called before forwarding the request to the upstream provider so the @@ -540,7 +551,7 @@ class ResponsesAPIRequestUtils: if isinstance(item, dict): item_id = item.get("id") if item_id and isinstance(item_id, str): - decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(item_id) + decoded = ResponsesAPIRequestUtils.decode_encrypted_item_id(item_id) if decoded: item["id"] = decoded["item_id"] @@ -549,12 +560,14 @@ class ResponsesAPIRequestUtils: ( _, unwrapped, - ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(encrypted_content) + ) = ResponsesAPIRequestUtils.unwrap_encrypted_content_with_model_id(encrypted_content) if unwrapped != encrypted_content: item["encrypted_content"] = unwrapped return request_input + _restore_encrypted_content_item_ids_in_input = restore_encrypted_content_item_ids_in_input + @staticmethod def strip_encrypted_reasoning_from_input( request_input: object, @@ -621,7 +634,7 @@ class ResponsesAPIRequestUtils: return f"resp_{base64_encoded_id}" @staticmethod - def _decode_responses_api_response_id( + def decode_responses_api_response_id( response_id: str, ) -> DecodedResponseId: """ @@ -673,9 +686,11 @@ class ResponsesAPIRequestUtils: response_id=response_id, ) + _decode_responses_api_response_id = decode_responses_api_response_id + @staticmethod def _is_litellm_encoded_response_id(response_id: str) -> bool: - decoded_response_id: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(response_id) + decoded_response_id: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(response_id) return ( decoded_response_id.get("model_id") is not None or decoded_response_id.get("custom_llm_provider") is not None @@ -686,7 +701,7 @@ class ResponsesAPIRequestUtils: """Get the model_id from the response_id""" if response_id is None: return None - decoded_response_id: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(response_id) + decoded_response_id: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(response_id) return decoded_response_id.get("model_id") or None @staticmethod @@ -706,11 +721,11 @@ class ResponsesAPIRequestUtils: Returns: The original previous_response_id """ - decoded_response_id: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(previous_response_id) + decoded_response_id: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(previous_response_id) return decoded_response_id.get("response_id", previous_response_id) @staticmethod - def _build_container_id( + def build_container_id( custom_llm_provider: str | None, model_id: str | None, container_id: str, @@ -726,8 +741,10 @@ class ResponsesAPIRequestUtils: base64_encoded_id: Final = base64.b64encode(assembled_id.encode("utf-8")).decode("utf-8") return f"cntr_{base64_encoded_id}" + _build_container_id = build_container_id + @staticmethod - def _decode_container_id(container_id: str) -> DecodedResponseId: + def decode_container_id(container_id: str) -> DecodedResponseId: """Decode a managed container ID to extract provider, model, and original container ID. Returns: @@ -786,6 +803,8 @@ class ResponsesAPIRequestUtils: response_id=container_id, ) + _decode_container_id = decode_container_id + @staticmethod def decode_container_id_to_original(container_id: str) -> str: """Decode a managed container ID to get the original provider-issued ID. @@ -793,11 +812,11 @@ class ResponsesAPIRequestUtils: This is used when making upstream API calls - we need to send the original container ID that the provider issued, not our encoded version. """ - decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) + decoded: Final = ResponsesAPIRequestUtils.decode_container_id(container_id) return decoded.get("response_id", container_id) @staticmethod - def _encode_container_ids_in_annotations( + def encode_container_ids_in_annotations( annotations: object, custom_llm_provider: str | None, model_id: str | None, @@ -806,12 +825,14 @@ class ResponsesAPIRequestUtils: if not annotations or not _is_object_sequence(annotations): return for ann in annotations: - ResponsesAPIRequestUtils._encode_container_id_on_output_item( + ResponsesAPIRequestUtils.encode_container_id_on_output_item( ann, custom_llm_provider, model_id, ) + _encode_container_ids_in_annotations = encode_container_ids_in_annotations + @staticmethod def _encode_container_ids_in_message_content( content: object, @@ -824,20 +845,20 @@ class ResponsesAPIRequestUtils: if _is_object_sequence(content): for part in content: if _is_object_dict(part): - ResponsesAPIRequestUtils._encode_container_ids_in_annotations( + ResponsesAPIRequestUtils.encode_container_ids_in_annotations( part.get("annotations"), custom_llm_provider, model_id, ) else: - ResponsesAPIRequestUtils._encode_container_ids_in_annotations( + ResponsesAPIRequestUtils.encode_container_ids_in_annotations( getattr(part, "annotations", None), custom_llm_provider, model_id, ) @staticmethod - def _encode_container_id_on_output_item( + def encode_container_id_on_output_item( item: object, custom_llm_provider: str | None, model_id: str | None, @@ -856,10 +877,10 @@ class ResponsesAPIRequestUtils: return def _maybe_encode(container_id: str) -> str | None: - decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) + decoded: Final = ResponsesAPIRequestUtils.decode_container_id(container_id) if decoded.get("custom_llm_provider") is not None: return None - return ResponsesAPIRequestUtils._build_container_id( + return ResponsesAPIRequestUtils.build_container_id( custom_llm_provider=custom_llm_provider, model_id=model_id, container_id=container_id, @@ -900,7 +921,7 @@ class ResponsesAPIRequestUtils: nested_obj: Final[object] = getattr(item, "code_interpreter_call", None) if nested_obj is not None: - ResponsesAPIRequestUtils._encode_container_id_on_output_item( + ResponsesAPIRequestUtils.encode_container_id_on_output_item( nested_obj, custom_llm_provider, model_id, @@ -913,6 +934,8 @@ class ResponsesAPIRequestUtils: model_id, ) + _encode_container_id_on_output_item = encode_container_id_on_output_item + @staticmethod def _collect_container_ids_from_annotations( annotations: object, @@ -1024,7 +1047,7 @@ class ResponsesAPIRequestUtils: return responses_api_response for item in output: - ResponsesAPIRequestUtils._encode_container_id_on_output_item( + ResponsesAPIRequestUtils.encode_container_id_on_output_item( item=item, custom_llm_provider=custom_llm_provider, model_id=model_id, @@ -1139,7 +1162,7 @@ class ResponsesAPIRequestUtils: class ResponseAPILoggingUtils: @staticmethod - def _is_response_api_usage(usage: dict | ResponseAPIUsage) -> bool: + def is_response_api_usage(usage: Mapping[str, object] | ResponseAPIUsage) -> bool: """returns True if usage is from OpenAI Response API""" if isinstance(usage, ResponseAPIUsage): return True @@ -1147,8 +1170,10 @@ class ResponseAPILoggingUtils: return True return False + _is_response_api_usage = is_response_api_usage + @staticmethod - def _transform_response_api_usage_to_chat_usage( + def transform_response_api_usage_to_chat_usage( usage_input: Mapping[str, object] | ResponseAPIUsage | Usage | None, ) -> Usage: """ @@ -1170,7 +1195,7 @@ class ResponseAPILoggingUtils: ) if isinstance(usage_input, Usage): return usage_input - if isinstance(usage_input, dict) and not ResponseAPILoggingUtils._is_response_api_usage(usage_input): + if isinstance(usage_input, dict) and not ResponseAPILoggingUtils.is_response_api_usage(usage_input): return Usage(**usage_input) response_api_usage: ResponseAPIUsage if isinstance(usage_input, dict): @@ -1261,3 +1286,5 @@ class ResponseAPILoggingUtils: setattr(chat_usage, "cost", response_api_usage.cost) return chat_usage + + _transform_response_api_usage_to_chat_usage = transform_response_api_usage_to_chat_usage diff --git a/litellm/router.py b/litellm/router.py index f93f12f2f2e..634fe2f3734 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -26,6 +26,7 @@ from collections.abc import ( AsyncIterator, Callable, Generator, + Iterable, Iterator, Mapping, MutableMapping, @@ -77,11 +78,11 @@ from litellm.integrations.otel.routing import routing_decision_attributes from litellm.integrations.otel.runtime import phase_attributes, phase_event, phase_span from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, coerce_token_limit, get_litellm_metadata_from_kwargs, get_metadata_variable_name_from_kwargs, get_or_create_metadata_bucket, + get_parent_otel_span_from_kwargs, ) from litellm.litellm_core_utils.coroutine_checker import coroutine_checker from litellm.litellm_core_utils.credential_accessor import CredentialAccessor @@ -142,13 +143,13 @@ from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage from litellm.router_strategy.simple_shuffle import simple_shuffle from litellm.router_strategy.tag_based_routing import ( - _get_tags_from_request_kwargs, get_deployments_for_tag, + get_tags_from_request_kwargs, is_valid_deployment_tag, ) from litellm.router_utils.access_windows import access_windows_config_error, filter_reserved_deployments from litellm.router_utils.add_retry_fallback_headers import ( - _HiddenParamsHost, + HiddenParamsHost, add_fallback_headers_to_response, add_retry_headers_to_response, apply_quality_router_decision_headers, @@ -171,7 +172,7 @@ from litellm.router_utils.auto_router_model_naming import ( count_capability_routers, ) from litellm.router_utils.batch_utils import ( - _get_router_metadata_variable_name, + get_router_metadata_variable_name, is_batch_retrieve_call_type, replace_model_in_jsonl, should_replace_model_in_jsonl, @@ -182,12 +183,12 @@ from litellm.router_utils.clientside_credential_handler import ( is_clientside_credential, ) from litellm.router_utils.common_utils import ( - _is_proxy_admin_request, filter_team_based_models, filter_web_search_deployments, format_fallback_outcome_message, format_no_fallback_group_message, get_request_team_id, + is_proxy_admin_request, provider_for_generic_call, resolve_model_group_alias, truncate_fallback_error_detail, @@ -196,22 +197,22 @@ from litellm.router_utils.common_utils import ( from litellm.router_utils.cooldown_cache import CooldownCache from litellm.router_utils.cooldown_handlers import ( DEFAULT_COOLDOWN_TIME_SECONDS, - _async_get_cooldown_deployments, - _async_get_cooldown_deployments_with_debug_info, _first_present, # pyright: ignore[reportPrivateUsage] - shared internal helper across router_utils submodules, matching the other cooldown_handlers imports on this line - _get_cooldown_deployments, - _set_cooldown_deployments, + async_get_cooldown_deployments, + async_get_cooldown_deployments_with_debug_info, + get_cooldown_deployments, is_advisor_orchestration_failure, is_background_response_cost_poll_not_found, is_caller_timeout_408, + set_cooldown_deployments, ) from litellm.router_utils.fallback_event_handlers import ( MID_STREAM_FALLBACK_CONTROLS_KEY, AttemptedFallbackTargets, - _check_non_standard_fallback_format, attempted_retries_for_request, carry_over_pre_routing_selection, carry_over_routed_deployment, + check_non_standard_fallback_format, clear_pre_routing_selection, committed_retry_budget_for_request, fallback_lookup_groups, @@ -623,7 +624,7 @@ def _anthropic_stream_fallback_error_for_raised( status_code: Final = _anthropic_stream_raised_error_status(error) if status_code is None: return _anthropic_stream_pre_content_error(error, model) - retriable: Final = litellm._should_retry(status_code) # pyright: ignore[reportPrivateUsage] # shared retry rule + retriable: Final = litellm.should_retry(status_code) return _anthropic_stream_pre_content_error(error, model) if retriable else None @@ -815,7 +816,7 @@ _live_routers: Final["weakref.WeakSet[Router]"] = weakref.WeakSet() def _replay_live_router_model_cost() -> None: """Re-assert every live router's deployments after the cost map is refreshed.""" for router in tuple(_live_routers): - router._replay_model_cost_registrations() + router.replay_model_cost_registrations() set_live_deployment_replay(_replay_live_router_model_cost) @@ -869,6 +870,14 @@ def as_output_cap(value: object) -> int | None: class Router: + @property + def _routing_groups(self) -> dict[str, RoutingGroup]: + return self.routing_groups + + @_routing_groups.setter + def _routing_groups(self, value: dict[str, RoutingGroup]) -> None: + self.routing_groups = value + model_names: set = set() cache_responses: bool | None = False default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour @@ -1404,7 +1413,7 @@ class Router: else: return RedisCache(**cache_config) - def _update_redis_cache(self, cache: RedisCache): + def update_redis_cache(self, cache: RedisCache) -> None: """ Update the redis cache for the router, if none set. @@ -1418,6 +1427,8 @@ class Router: self.cache.attach_redis_cache(cache) self._claude_code_session_router_cache.attach_redis_cache(cache) + _update_redis_cache = update_redis_cache + # Maps a routing strategy string to the attribute on `self` that holds # the default group's strategy selector for that strategy. (The selectors # double as `CustomLogger` callbacks, hence the legacy `*_logger` attrs.) @@ -1638,7 +1649,7 @@ class Router: if selector is not None: self._register_router_selector(selector) - self._routing_groups: dict[str, RoutingGroup] = {group.group_name: group for group, _ in built} + self.routing_groups: dict[str, RoutingGroup] = {group.group_name: group for group, _ in built} self._model_to_group: dict[str, str] = { model_name: group.group_name for group, _ in built @@ -1662,9 +1673,9 @@ class Router: models win over indirection); config-time collisions are rejected by `_init_routing_groups`. """ - if not self._routing_groups: + if not self.routing_groups: return None - group: Final = self._routing_groups.get(model_name) + group: Final = self.routing_groups.get(model_name) if ( group is None or model_name in self.model_name_to_deployment_indices @@ -1687,7 +1698,7 @@ class Router: specific deployment > model id > model_group_alias > routing group > model_name > team/pattern/default fallbacks. """ - if not self._routing_groups: + if not self.routing_groups: return None routing_group: Final = self.get_routing_group(model) if routing_group is None: @@ -1695,11 +1706,11 @@ class Router: return [ apply_routing_group_priority(routing_group, member, deployment) for member in routing_group.models - for deployment in self._get_all_deployments(model_name=member, team_id=team_id) + for deployment in self.get_all_deployments(model_name=member, team_id=team_id) ] def _is_priority_routing_group(self, model: str) -> bool: - resolved: Final = self._get_model_from_alias(model=model) or model + resolved: Final = self.get_model_from_alias(model=model) or model group: Final = self.get_routing_group(resolved) return group is not None and group.routing_strategy == "priority" @@ -1729,7 +1740,7 @@ class Router: """ if model_group is None: return False - resolved: Final = self._get_model_from_alias(model=model_group) or model_group + resolved: Final = self.get_model_from_alias(model=model_group) or model_group group: Final = self.get_routing_group(resolved) if group is None: return False @@ -1831,7 +1842,7 @@ class Router: def _globally_registered_strategies(self) -> frozenset[str]: configured: Final = ( self.routing_strategy, - *(group.routing_strategy for group in self._routing_groups.values()), + *(group.routing_strategy for group in self.routing_groups.values()), ) return frozenset( normalized for normalized in map(self._normalize_strategy, configured) if normalized is not None @@ -1885,7 +1896,7 @@ class Router: self._bind_override_selector_to_request(override, override_selector, request_kwargs) return override, override_selector - resolved_model: Final = self._get_model_from_alias(model=model) or model + resolved_model: Final = self.get_model_from_alias(model=model) or model group_name: Final = ( resolved_model if self.get_routing_group(resolved_model) is not None @@ -1898,7 +1909,7 @@ class Router: verbose_router_logger.debug("routing_group=default model=%s strategy=%s", model, strategy) return strategy, selector - group: Final = self._routing_groups[group_name] + group: Final = self.routing_groups[group_name] if group.routing_strategy == "priority": return "simple-shuffle", None strategy = self._normalize_strategy(group.routing_strategy) @@ -2493,7 +2504,7 @@ class Router: cache=self.cache, is_priority_group=self._is_priority_routing_group ) elif pre_call_check == "router_budget_limiting": - if self._get_router_deployment_budget_limiter() is not None: + if self.get_router_deployment_budget_limiter() is not None: continue _callback = RouterBudgetLimiting( dual_cache=self.cache, @@ -2916,7 +2927,7 @@ class Router: call_type="acompletion", start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) @@ -2987,7 +2998,7 @@ class Router: if not isinstance(item_headers, dict): item_headers = {} - cast(_HiddenParamsHost, fallback_item)._hidden_params = { + cast(HiddenParamsHost, fallback_item)._hidden_params = { **item_hidden_params, **fallback_hidden_params, "additional_headers": {**item_headers, **fallback_headers}, @@ -3364,7 +3375,7 @@ class Router: from litellm.exceptions import MidStreamFallbackError from litellm.responses.streaming_iterator import ( BaseResponsesAPIStreamingIterator, - _get_openai_response_types, + get_openai_response_types, ) source_iterator: Final = response @@ -3373,7 +3384,7 @@ class Router: # per-chunk type check inside FallbackResponsesStreamWrapper # stays cheap; mirrors the source-iterator filter at # responses/streaming_iterator.py:243-247. - _openai_types: Final = _get_openai_response_types() + _openai_types: Final = get_openai_response_types() _RESPONSES_TERMINAL_EVENT_TYPES: Final = ( _openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED, _openai_types.ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, @@ -3486,7 +3497,9 @@ class Router: async def stream_with_fallbacks(): held_lifecycle_events: tuple[object, ...] = () # rebind-ok: flushed at first output, dropped on fallback try: - async for item in source_iterator: + async for item in cast( # cast-ok: response iterators are async streams + AsyncIterator[object], source_iterator + ): if _responses_stream_holds_event(item, len(held_lifecycle_events)): held_lifecycle_events = (*held_lifecycle_events, item) continue @@ -3701,7 +3714,7 @@ class Router: wrapper_ref, fallback_response ) fallback_headers_are_settled = False - for fallback_item in fallback_response: + for fallback_item in cast(Iterable[object], fallback_response): # cast-ok: __iter__ was checked if not fallback_headers_are_settled: fallback_headers_are_settled = True # a fallback that failed over again only repoints itself once it yields @@ -3788,7 +3801,7 @@ class Router: input_kwargs_for_streaming_fallback: Final = kwargs.copy() input_kwargs_for_streaming_fallback["model"] = model - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) start_time: Final = time.time() deployment = await self.async_get_available_deployment( model=model, @@ -3812,7 +3825,7 @@ class Router: call_type="async_get_available_deployment", start_time=start_time, end_time=end_time, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) ) @@ -3955,7 +3968,7 @@ class Router: """ kwargs.setdefault("litellm_trace_id", str(uuid.uuid4())) model_group_alias: str | None = None - if self._get_model_from_alias(model=model): + if self.get_model_from_alias(model=model): model_group_alias = model kwargs.setdefault(metadata_variable_name, {}).update( {"model_group": model, "model_group_alias": model_group_alias} @@ -4056,7 +4069,7 @@ class Router: litellm_params: Final = deployment["litellm_params"].copy() dynamic_litellm_params: Final = get_dynamic_litellm_params(litellm_params=litellm_params, request_kwargs=kwargs) # Use deployment model_name as model_group for generating model_id - metadata_variable_name: Final = _get_router_metadata_variable_name( + metadata_variable_name: Final = get_router_metadata_variable_name( function_name=function_name, ) model_group: Final = kwargs.get(metadata_variable_name, {}).get("model_group") @@ -4122,7 +4135,7 @@ class Router: deployment_litellm_model_name = deployment_pydantic_obj.litellm_params.model deployment_api_base = deployment_pydantic_obj.litellm_params.api_base - metadata_variable_name: Final = _get_router_metadata_variable_name( + metadata_variable_name: Final = get_router_metadata_variable_name( function_name=function_name, ) @@ -4144,7 +4157,7 @@ class Router: kwargs[metadata_variable_name].setdefault( ROUTING_REQUEST_TAGS_METADATA_KEY, - tuple(_get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name)), + tuple(get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name)), ) ## DEPLOYMENT-LEVEL TAGS @@ -4503,7 +4516,7 @@ class Router: stream=False, **kwargs, ): - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) ### FLOW ITEM ### _request_id: Final = str(uuid.uuid4()) item: Final = FlowItem( @@ -4563,7 +4576,7 @@ class Router: args: tuple[object, ...], kwargs: dict[str, object], ): - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) ### FLOW ITEM ### _request_id: Final = str(uuid.uuid4()) item: Final = FlowItem( @@ -4797,7 +4810,7 @@ class Router: model_name = model try: verbose_router_logger.debug("Inside _image_generation()- model: %s; kwargs: %s", model, kwargs) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "prompt"}], @@ -4882,7 +4895,7 @@ class Router: model_name: Final = model try: verbose_router_logger.debug("Inside _atranscription()- model: %s; kwargs: %s", model, kwargs) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "prompt"}], @@ -4975,7 +4988,7 @@ class Router: model_name: Final = model try: verbose_router_logger.debug("Inside _aspeech()- model: %s; kwargs: %s", model, kwargs) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "prompt"}], @@ -5147,7 +5160,7 @@ class Router: async def _atext_completion(self, model: str, prompt: str, **kwargs): try: verbose_router_logger.debug("Inside _atext_completion()- model: %s; kwargs: %s", model, kwargs) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": prompt}], @@ -5217,7 +5230,7 @@ class Router: async def _aadapter_completion(self, adapter_id: str, model: str, **kwargs): try: verbose_router_logger.debug("Inside _aadapter_completion()- model: %s; kwargs: %s", model, kwargs) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "default text"}], @@ -5366,7 +5379,7 @@ class Router: # Use simple_shuffle for weighted selection return simple_shuffle( - resolve_model_alias=self._get_model_from_alias, + resolve_model_alias=self.get_model_from_alias, healthy_deployments=healthy_deployments, model=guardrail_name, request_kwargs=None, @@ -5443,7 +5456,7 @@ class Router: function_name: Final = "_ageneric_api_call_with_fallbacks" deployment = None # rebind-ok: pre-init so the except block can stamp a failure with no deployment picked try: - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) try: deployment = await self.async_get_available_deployment( # rebind-ok: set on success, see pre-init above model=model, @@ -5907,7 +5920,7 @@ class Router: model_group: Final = cast(str, initial_kwargs.get("model")) # cast-ok: kwargs always carries the model group retry_kwargs: Final = mid_stream_retry_kwargs(initial_kwargs) healthy_deployments, all_deployments = await self._async_get_healthy_deployments( - model=model_group, parent_otel_span=_get_parent_otel_span_from_kwargs(retry_kwargs) + model=model_group, parent_otel_span=get_parent_otel_span_from_kwargs(retry_kwargs) ) budget, policy_applies = self._anthropic_messages_retry_budget(_mid_stream_retry_trigger(e), initial_kwargs) last_error = e # rebind-ok: the newest failure is what the fallback chain and the caller see @@ -6121,7 +6134,7 @@ class Router: The response from the handler function """ handler_name: Final = original_function.__name__ - metadata_variable_name: Final = _get_router_metadata_variable_name(function_name="generic_api_call") + metadata_variable_name: Final = get_router_metadata_variable_name(function_name="generic_api_call") try: verbose_router_logger.debug( "Inside _generic_api_call() - handler: %s, model: %s; kwargs: %s", handler_name, model, kwargs @@ -6270,7 +6283,7 @@ class Router: model_name = None try: verbose_router_logger.debug("Inside _aembedding()- model: %s; kwargs: %s", model, kwargs) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, input=input, @@ -6340,7 +6353,7 @@ class Router: from litellm.router_utils.common_utils import add_model_file_id_mappings verbose_router_logger.debug("Inside _atext_completion()- model: %s; kwargs: %s", model, kwargs) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) healthy_deployments: Final = await self.async_get_healthy_deployments( model=model, messages=[{"role": "user", "content": "files-api-fake-text"}], @@ -6477,7 +6490,7 @@ class Router: from litellm.vector_stores import acreate as avector_store_create_sdk - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "vector-store-api-fake-text"}], @@ -6548,7 +6561,7 @@ class Router: try: kwargs["model"] = model kwargs["original_function"] = self._acreate_batch - metadata_variable_name: Final = _get_router_metadata_variable_name(function_name="_acreate_batch") + metadata_variable_name: Final = get_router_metadata_variable_name(function_name="_acreate_batch") self._update_kwargs_before_fallbacks( model=model, kwargs=kwargs, @@ -6575,7 +6588,7 @@ class Router: ) -> LiteLLMBatch: try: verbose_router_logger.debug("Inside _acreate_batch()- model: %s; kwargs: %s", model, kwargs) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "files-api-fake-text"}], @@ -6636,9 +6649,9 @@ class Router: Future Improvement - cache the result. """ try: - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) requested_model_group: Final = model - metadata_variable_name: Final = _get_router_metadata_variable_name(function_name="aretrieve_batch") + metadata_variable_name: Final = get_router_metadata_variable_name(function_name="aretrieve_batch") if model is not None: filtered_model_list: ( list[DeploymentTypedDict] | list[dict] | dict | None @@ -6746,7 +6759,7 @@ class Router: try: kwargs["model"] = model kwargs["original_function"] = self._acancel_batch - metadata_variable_name: Final = _get_router_metadata_variable_name(function_name="_acancel_batch") + metadata_variable_name: Final = get_router_metadata_variable_name(function_name="_acancel_batch") self._update_kwargs_before_fallbacks( model=model, kwargs=kwargs, @@ -6773,7 +6786,7 @@ class Router: ) -> LiteLLMBatch: try: verbose_router_logger.debug("Inside _acancel_batch()- model: %s; kwargs: %s", model, kwargs) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) deployment: Final = await self.async_get_available_deployment( model=model, messages=[{"role": "user", "content": "batch-api-fake-text"}], @@ -7383,7 +7396,7 @@ class Router: container_id: Final = kwargs.get("container_id") _forwarded_model_id: Final = kwargs.get("model_id") if isinstance(container_id, str): - decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) + decoded: Final = ResponsesAPIRequestUtils.decode_container_id(container_id) original_id: Final = decoded.get("response_id", container_id) if original_id != container_id: kwargs["container_id"] = original_id @@ -7540,7 +7553,7 @@ class Router: # that fails with RouterRateLimitError whenever the "remaining" entries # are all in cooldown — the inner async_get_healthy_deployments call # would find an empty list and raise immediately. - cooldown_ids = set(await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=None)) + cooldown_ids = set(await async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=None)) remaining: Final = (all_ids - cooldown_ids) - excluded if not remaining: return None @@ -7635,9 +7648,9 @@ class Router: order_model_group: Final = get_pre_routing_selection(kwargs) or original_model_group all_deployments: Final = self.get_model_list(model_name=order_model_group, team_id=_request_team_id) or () _order_set: Final[set] = { - litellm.utils._get_deployment_order(d) + litellm.utils.get_deployment_order(d) for d in all_deployments - if litellm.utils._get_deployment_order(d) is not None + if litellm.utils.get_deployment_order(d) is not None } order_values: Final[list] = sorted(_order_set) if len(order_values) > 1 and not _skip_order_fallback: @@ -7651,7 +7664,7 @@ class Router: # Get external fallbacks — handle both standard and non-standard formats external_fallback_group: list | None = None if fallbacks is not None and lookup_groups: - if _check_non_standard_fallback_format(fallbacks=fallbacks): + if check_non_standard_fallback_format(fallbacks=fallbacks): # Non-standard formats (e.g. ["claude-3-haiku"] or # [{"model": "...", "messages": [...]}]) are passed through directly external_fallback_group = fallbacks @@ -7696,7 +7709,7 @@ class Router: verbose_router_logger.info("Trying to fallback b/w models") # check if client-side fallbacks are used (e.g. fallbacks = ["gpt-3.5-turbo", "claude-3-haiku"] or fallbacks=[{"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hey, how's it going?"}]}] - is_non_standard_fallback_format: Final = _check_non_standard_fallback_format(fallbacks=fallbacks) + is_non_standard_fallback_format: Final = check_non_standard_fallback_format(fallbacks=fallbacks) if is_non_standard_fallback_format: input_kwargs.update( @@ -7821,9 +7834,9 @@ class Router: return response except Exception as new_exception: - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) fallback_failure_exception_str = truncate_fallback_error_detail(redact_string(str(new_exception))) - cooldown_info: Final = await _async_get_cooldown_deployments_with_debug_info( + cooldown_info: Final = await async_get_cooldown_deployments_with_debug_info( litellm_router_instance=self, parent_otel_span=parent_otel_span, ) @@ -7862,7 +7875,7 @@ class Router: kwargs["_context_compaction_state"] = initialize_compaction_state(kwargs, compaction_surface) clear_pre_routing_selection(kwargs) # pyright: ignore[reportUnknownArgumentType] # **kwargs is untyped at this boundary if not isinstance(kwargs.get("attempted_targets"), AttemptedFallbackTargets): - _fallback_metadata_key: Final = _get_router_metadata_variable_name( + _fallback_metadata_key: Final = get_router_metadata_variable_name( function_name=getattr(kwargs.get("original_function"), "__name__", None) ) _sibling_metadata_key: Final = ( @@ -7974,7 +7987,7 @@ class Router: status_code: Final = getattr(exception, "status_code", None) if not failed_deployment_id or not isinstance(status_code, int): return () - if litellm._should_retry(status_code): # pyright: ignore[reportPrivateUsage] # as in should_retry_this_error + if litellm.should_retry(status_code): return () already_skipped_ids: Final = _as_retry_skipped_deployment_ids(already_skipped) skipped: Final = tuple(sorted(frozenset((*already_skipped_ids, failed_deployment_id)))) @@ -7988,7 +8001,7 @@ class Router: verbose_router_logger.debug("Inside async function with retries.") original_function: Final = kwargs.pop("original_function") fallbacks: Final = kwargs.pop("fallbacks", self.fallbacks) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) context_window_fallbacks: Final = kwargs.pop("context_window_fallbacks", self.context_window_fallbacks) content_policy_fallbacks: Final = kwargs.pop("content_policy_fallbacks", self.content_policy_fallbacks) # Support per-request model_group_retry_policy override (from key/team settings) @@ -8209,7 +8222,7 @@ class Router: num_retries: int | None = None if available_models is not None and len(available_models) == 1: - num_retries = cast(int | None, available_models[0]["litellm_params"].get("num_retries")) + num_retries = available_models[0]["litellm_params"].get("num_retries") if mock_testing_rate_limit_error is not None and mock_testing_rate_limit_error is True: verbose_router_logger.info( @@ -8255,7 +8268,7 @@ class Router: raise error status_code: Final = getattr(error, "status_code", None) - if status_code is not None and not litellm._should_retry(status_code): + if status_code is not None and not litellm.should_retry(status_code): # 401/403 are special cases - allow retry if multiple deployments exist (handled below) if status_code not in (401, 403): raise error @@ -8346,7 +8359,7 @@ class Router: response_headers = e.litellm_response_headers if response_headers is not None: - timeout = litellm._calculate_retry_after( + timeout = litellm.calculate_retry_after( remaining_retries=remaining_retries, max_retries=num_retries, response_headers=response_headers, @@ -8354,7 +8367,7 @@ class Router: ) else: - timeout = litellm._calculate_retry_after( + timeout = litellm.calculate_retry_after( remaining_retries=remaining_retries, max_retries=num_retries, min_timeout=self.retry_after, @@ -8408,7 +8421,7 @@ class Router: model_group=model_group, total_tokens=total_tokens if counted_tokens is None else max(0, total_tokens - counted_tokens), rpm_increment=1 if counted_tokens is None else 0, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) except Exception as e: @@ -8442,7 +8455,7 @@ class Router: model_group=model_group, total_tokens=total_tokens, rpm_increment=1, - parent_otel_span=_get_parent_otel_span_from_kwargs(request_kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(request_kwargs), ) except Exception: deployment_metadata.pop(ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY, None) @@ -8631,7 +8644,7 @@ class Router: litellm_router_instance=self, deployment_id=deployment_id, ) - result: Final = _set_cooldown_deployments( + result: Final = set_cooldown_deployments( litellm_router_instance=self, exception_status=exception_status, original_exception=exception, @@ -8668,7 +8681,7 @@ class Router: return elif isinstance(id, int): id = str(id) - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) dt: Final = get_utc_datetime() current_minute: Final = dt.strftime("%H-%M") # use the same timezone regardless of system clock @@ -8830,9 +8843,9 @@ class Router: return tuple( sorted( { - litellm.utils._get_deployment_order(d) + litellm.utils.get_deployment_order(d) for d in all_deployments - if litellm.utils._get_deployment_order(d) is not None + if litellm.utils.get_deployment_order(d) is not None } ) ) @@ -8861,7 +8874,7 @@ class Router: fallbacks: Final = kwargs.get("fallbacks", self.fallbacks) if not fallbacks: return False - if _check_non_standard_fallback_format(fallbacks=fallbacks): + if check_non_standard_fallback_format(fallbacks=fallbacks): return True resolved, _ = get_fallback_model_group_for_lookup_groups( fallbacks=fallbacks, @@ -8913,7 +8926,7 @@ class Router: except Exception: pass - unhealthy_deployments: Final = _get_cooldown_deployments( + unhealthy_deployments: Final = get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) unhealthy_set: Final = set(unhealthy_deployments) @@ -8941,7 +8954,7 @@ class Router: except Exception: pass - unhealthy_deployments: Final = await _async_get_cooldown_deployments( + unhealthy_deployments: Final = await async_get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) # Convert to set for O(1) lookup instead of O(n) @@ -8999,7 +9012,7 @@ class Router: prefer_async_handlers=True, ) ) - _set_cooldown_deployments( + set_cooldown_deployments( litellm_router_instance=self, exception_status=e.status_code, original_exception=e, @@ -10488,7 +10501,7 @@ class Router: persist_across_reloads=False, ) - def _replay_model_cost_registrations(self) -> None: + def replay_model_cost_registrations(self) -> None: """Re-assert this router's deployments onto a freshly fetched catalog. Reads ``model_list`` at call time, so only deployments the router still @@ -10514,6 +10527,8 @@ class Router: ) self._invalidate_model_group_info_cache() + _replay_model_cost_registrations = replay_model_cost_registrations + def delete_deployment(self, id: str) -> Deployment | None: """ Parameters: @@ -10534,7 +10549,7 @@ class Router: self._invalidate_model_group_info_cache() self._invalidate_access_groups_cache() self._update_deployment_indices_after_removal(model_id=id, removal_idx=deployment_idx) - _budget_limiter: Final = self._get_router_deployment_budget_limiter() + _budget_limiter: Final = self.get_router_deployment_budget_limiter() if _budget_limiter is not None: _budget_limiter.unregister_deployment_budget(model_id=id) try: @@ -10553,7 +10568,7 @@ class Router: except Exception: return None - def _get_router_deployment_budget_limiter( + def get_router_deployment_budget_limiter( self, ) -> RouterBudgetLimiting | None: """ @@ -10572,6 +10587,8 @@ class Router: return _cb return None + _get_router_deployment_budget_limiter = get_router_deployment_budget_limiter + def _deployment_has_budget_limits(self, deployment: Deployment) -> bool: return ( deployment.litellm_params.get("max_budget") is not None @@ -10584,7 +10601,7 @@ class Router: if model_id is None: return - _budget_limiter = self._get_router_deployment_budget_limiter() + _budget_limiter = self.get_router_deployment_budget_limiter() if not self._deployment_has_budget_limits(deployment=deployment): if _budget_limiter is not None: @@ -10593,7 +10610,7 @@ class Router: if _budget_limiter is None: self.add_optional_pre_call_checks(optional_pre_call_checks=["router_budget_limiting"]) - _budget_limiter = self._get_router_deployment_budget_limiter() + _budget_limiter = self.get_router_deployment_budget_limiter() if _budget_limiter is not None: _budget_limiter.register_deployment_budget(deployment=deployment.to_json(exclude_none=True)) @@ -10626,7 +10643,7 @@ class Router: using a paused deployment. """ deployment: Final = self.get_deployment(model_id=model_id) - if deployment is None or self._is_deployment_blocked(deployment): + if deployment is None or self.is_deployment_blocked(deployment): return None return CredentialLiteLLMParams.model_validate( deployment.litellm_params.model_dump(exclude_none=True) @@ -10655,15 +10672,17 @@ class Router: return None @staticmethod - def _deployment_usable_by_team(model: Mapping | Deployment, team_id: str | None) -> bool: + def deployment_usable_by_team(model: Mapping[str, object] | Deployment, team_id: str | None) -> bool: """ A team-scoped deployment (``model_info.team_id`` set) is only usable by callers from that same team; deployments without a team owner are shared. """ model_info: Final = model.get("model_info") if isinstance(model, dict) else model.model_info - owner_team_id: Final = model_info.get("team_id") if model_info is not None else None + owner_team_id: Final = cast(Mapping[str, object], model_info).get("team_id") if model_info is not None else None return owner_team_id is None or owner_team_id == team_id + _deployment_usable_by_team = deployment_usable_by_team + def _get_model_group_deployment_usable_by_team( self, model_group_name: str, team_id: str | None ) -> Deployment | None: @@ -10674,7 +10693,7 @@ class Router: """ indices: Final = self.model_name_to_deployment_indices.get(model_group_name) or () usable: Final = ( - self.model_list[idx] for idx in indices if self._deployment_usable_by_team(self.model_list[idx], team_id) + self.model_list[idx] for idx in indices if self.deployment_usable_by_team(self.model_list[idx], team_id) ) first_usable: Final = next(usable, None) if first_usable is None: @@ -10886,7 +10905,7 @@ class Router: The model group a request to model_name routes to: its target when model_name is a `model_group_alias`, else model_name itself. """ - target: Final = self._get_model_from_alias(model_name) + target: Final = self.get_model_from_alias(model_name) return target if target is not None else model_name def _routable_deployments(self, model_name: str, team_id: str | None) -> tuple[DeploymentTypedDict, ...]: @@ -10902,10 +10921,8 @@ class Router: Returns an empty tuple for wildcard-expanded or unknown names. """ - named: Final = self._get_all_deployments(model_name=self.routable_model_group(model_name), team_id=team_id) - usable: Final = tuple( - deployment for deployment in named if self._deployment_usable_by_team(deployment, team_id) - ) + named: Final = self.get_all_deployments(model_name=self.routable_model_group(model_name), team_id=team_id) + usable: Final = tuple(deployment for deployment in named if self.deployment_usable_by_team(deployment, team_id)) return tuple( deployment for deployment in (usable or named) @@ -10956,7 +10973,7 @@ class Router: or self._get_team_public_name_deployment(model_id=model_id, team_id=team_id) or self._get_wildcard_deployment_usable_by_team(model_id=model_id, team_id=team_id) ) - if deployment is None or self._is_deployment_blocked(deployment): + if deployment is None or self.is_deployment_blocked(deployment): return None return deployment @@ -10975,7 +10992,7 @@ class Router: global_wildcard_models: Final = tuple( wildcard_model for wildcard_model in (self.pattern_router.route(model_id) or ()) - if self._deployment_usable_by_team(wildcard_model, team_id) + if self.deployment_usable_by_team(wildcard_model, team_id) ) potential_wildcard_models: Final = team_wildcard_models or global_wildcard_models if not potential_wildcard_models: @@ -11226,7 +11243,7 @@ class Router: 2. If not, check if litellm model name is in model info 3. If not, return None """ - from litellm.utils import _update_dictionary, cost_map_omits_token_price + from litellm.utils import cost_map_omits_token_price, update_dictionary model_info: ModelInfo | None = None custom_model_info: dict | None = None @@ -11260,7 +11277,7 @@ class Router: if base_model_info is not None: base_model_key = base_model_info.get("key") # Base model provides defaults, custom model info overrides - custom_model_info = _update_dictionary( + custom_model_info = update_dictionary( cast(dict, base_model_info), custom_model_info, ) @@ -11273,7 +11290,7 @@ class Router: # merge with custom overriding built-in model_info = cast( ModelInfo, - _update_dictionary( + update_dictionary( copy.deepcopy(cast(dict, litellm_model_name_model_info)), custom_model_info, ), @@ -11878,7 +11895,7 @@ class Router: check tell a genuine cross-group route from same-group unavailability without leaking deployment ids into request kwargs bound for the provider. """ - resolved: Final = self._get_model_from_alias(model=model) or model + resolved: Final = self.get_model_from_alias(model=model) or model routing_group_members: Final = self._get_routing_group_deployments(model=resolved, team_id=team_id) if routing_group_members is not None: return self._deployment_ids(routing_group_members) @@ -11890,7 +11907,7 @@ class Router: return self._deployment_ids( (early_deployments,) if isinstance(early_deployments, Mapping) else early_deployments ) - return self._deployment_ids(self._get_all_deployments(model_name=resolved, team_id=team_id)) + return self._deployment_ids(self.get_all_deployments(model_name=resolved, team_id=team_id)) @staticmethod def _deployment_ids(deployments: Sequence[Mapping[str, object]]) -> frozenset[str]: @@ -12027,7 +12044,7 @@ class Router: # No match: deployment is for a different team or doesn't match the requested model return False - def _get_all_deployments( + def get_all_deployments( self, model_name: str, model_alias: str | None = None, @@ -12101,6 +12118,8 @@ class Router: return returned_models + _get_all_deployments = get_all_deployments + def get_model_names(self, team_id: str | None = None) -> list[str]: """ Returns all possible model names for the router, including models defined via model_group_alias. @@ -12144,16 +12163,18 @@ class Router: return {name for name, fully_blocked in blocked_by_name.items() if fully_blocked} @staticmethod - def _are_all_deployments_blocked( + def are_all_deployments_blocked( deployments: list[DeploymentTypedDict], ) -> bool: return len(deployments) > 0 and all( (deployment.get("model_info") or {}).get("blocked") is True for deployment in deployments ) + _are_all_deployments_blocked = are_all_deployments_blocked + def _is_model_fully_blocked(self, model: str) -> bool: deployments: Final = self.get_model_list(model_name=model) or [] - return self._are_all_deployments_blocked(deployments=deployments) + return self.are_all_deployments_blocked(deployments=deployments) async def async_get_fully_unhealthy_model_names(self) -> set[str]: """ @@ -12269,9 +12290,7 @@ class Router: {**row, "model_name": model_alias} for row in self._materialize_routing_group_rows((alias_group,)) ) else: - returned_models.extend( - self._get_all_deployments(model_name=_router_model_name, model_alias=model_alias) - ) + returned_models.extend(self.get_all_deployments(model_name=_router_model_name, model_alias=model_alias)) return returned_models @@ -12279,7 +12298,7 @@ class Router: """ Callable routing groups materialized as model-list rows, mirroring `get_model_list_from_model_alias`: each member deployment is emitted - under the group's name (via `_get_all_deployments`' `model_alias` + under the group's name (via `get_all_deployments`' `model_alias` rewrite), which is what surfaces groups in `get_model_names`, `/v1/models` discovery, `get_model_group_usage`, and the blocked/unhealthy hiding that all read `get_model_list`. @@ -12293,7 +12312,7 @@ class Router: rows: Final = self._materialize_routing_group_rows( tuple( callable_group - for name in self._routing_groups + for name in self.routing_groups if (callable_group := self.get_routing_group(name)) is not None ) ) @@ -12305,7 +12324,7 @@ class Router: self._as_routing_group_row(apply_routing_group_priority(group, member, deployment)) for group in groups for member in group.models - for deployment in self._get_all_deployments(model_name=member, model_alias=group.group_name) + for deployment in self.get_all_deployments(model_name=member, model_alias=group.group_name) ) @staticmethod @@ -12440,7 +12459,7 @@ class Router: returned_models: list[DeploymentTypedDict] = [] if model_name is not None: - returned_models.extend(self._get_all_deployments(model_name=model_name, team_id=team_id)) + returned_models.extend(self.get_all_deployments(model_name=model_name, team_id=team_id)) returned_models.extend(self.get_model_list_from_model_alias(model_name=model_name)) returned_models.extend(self.get_model_list_from_routing_groups(model_name=model_name)) @@ -12568,7 +12587,7 @@ class Router: return access_groups - def _is_model_access_group_for_wildcard_route(self, model_access_group: str) -> bool: + def is_model_access_group_for_wildcard_route(self, model_access_group: str) -> bool: """ Return True if model access group is a wildcard route """ @@ -12587,6 +12606,8 @@ class Router: return False + _is_model_access_group_for_wildcard_route = is_model_access_group_for_wildcard_route + def get_settings(self): """ Get router settings method, returns a dictionary of the settings and their values. @@ -12625,7 +12646,7 @@ class Router: _settings_to_return["routing_groups"] = [ group.model_dump(exclude=frozenset(("model_priorities",)) if group.model_priorities is None else None) - for group in self._routing_groups.values() + for group in self.routing_groups.values() ] return _settings_to_return @@ -12709,7 +12730,7 @@ class Router: The appropriate client based on the given client_type and kwargs. """ model_id: Final = deployment["model_info"]["id"] - parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(kwargs) if client_type == "max_parallel_requests": cache_key = f"{model_id}_max_parallel_requests_client" client = self.cache.get_cache(key=cache_key, local_only=True, parent_otel_span=parent_otel_span) @@ -12884,7 +12905,7 @@ class Router: _context_window_error = False _potential_error_str = "" _rate_limit_error = False - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(request_kwargs) has_countable_input: Final = (messages is not None or input is not None) and not compaction_pending( request_kwargs @@ -13030,7 +13051,7 @@ class Router: return _returned_deployments - def _get_model_from_alias(self, model: str) -> str | None: + def get_model_from_alias(self, model: str) -> str | None: """ Get the model from the alias. @@ -13040,6 +13061,8 @@ class Router: """ return resolve_model_group_alias(self.model_group_alias, model) + _get_model_from_alias = get_model_from_alias + def _get_deployment_by_litellm_model(self, model: str) -> list: """ Get the deployment by litellm model. @@ -13062,7 +13085,7 @@ class Router: # This intentionally takes priority over team pattern routers below, # so that named team deployments shadow wildcard/pattern routes. if request_team_id is not None: - team_deployments = self._get_all_deployments(model_name=model, team_id=request_team_id) + team_deployments = self.get_all_deployments(model_name=model, team_id=request_team_id) if team_deployments: return model, team_deployments elif include_team_models: @@ -13125,10 +13148,10 @@ class Router: global, then admin-across-teams resolution `_common_checks_available_deployment` applies, so strategy selection and compression policy can never disagree with deployment selection about which marker a name means.""" - registered_name: Final = self._get_model_from_alias(model=model) or model + registered_name: Final = self.get_model_from_alias(model=model) or model team_id: Final = get_request_team_id(request_kwargs) - deployments: Final = self._get_all_deployments(model_name=registered_name, team_id=team_id) - if deployments or team_id is not None or not _is_proxy_admin_request(request_kwargs): + deployments: Final = self.get_all_deployments(model_name=registered_name, team_id=team_id) + if deployments or team_id is not None or not is_proxy_admin_request(request_kwargs): return deployments return self._team_deployments_across_teams(registered_name) @@ -13190,7 +13213,7 @@ class Router: f"LiteLLM Router: Trying to call specific deployment, but Model ID :{model} does not exist in Model ID map" ) - _model_from_alias: Final = self._get_model_from_alias(model=model) + _model_from_alias: Final = self.get_model_from_alias(model=model) if _model_from_alias is not None: model = _model_from_alias @@ -13199,7 +13222,7 @@ class Router: early: Final = self._try_early_resolve_deployments_for_model_not_in_names( model=model, request_team_id=request_team_id, - include_team_models=_is_proxy_admin_request(request_kwargs), + include_team_models=is_proxy_admin_request(request_kwargs), ) if early is not None: if not isinstance(early[1], list): @@ -13229,7 +13252,7 @@ class Router: healthy_deployments = ( _routing_group_deployments if _routing_group_deployments is not None - else self._get_all_deployments(model_name=model, team_id=request_team_id) + else self.get_all_deployments(model_name=model, team_id=request_team_id) ) _pre_model_access_group_filter_len: Final = len(healthy_deployments) healthy_deployments = self._filter_reserved_deployments( @@ -13286,7 +13309,7 @@ class Router: ) # Re-assign model to the fallback and try to get deployments again model = fallback_model - healthy_deployments = self._get_all_deployments(model_name=model, team_id=request_team_id) + healthy_deployments = self.get_all_deployments(model_name=model, team_id=request_team_id) healthy_deployments = self._filter_reserved_deployments( model=model, healthy_deployments=self._filter_deployments_by_model_access_groups( @@ -13461,7 +13484,7 @@ class Router: routing_read_batch: Final = RoutingReadBatch.active() cooldown_deployments: Final = ( - await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) + await async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) if routing_read_batch is None else await routing_read_batch.async_get_cooldown_deployments( litellm_router_instance=self, @@ -13496,7 +13519,7 @@ class Router: ) if self.enable_pre_call_checks and (messages is not None or input is not None): - deployments_to_check: Final = cast(list[dict], healthy_deployments) + deployments_to_check: Final = healthy_deployments healthy_deployments = self._pre_call_checks( model=model, healthy_deployments=deployments_to_check, @@ -13730,7 +13753,7 @@ class Router: request_kwargs=request_kwargs, ) try: - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(request_kwargs) ######################################################### # Execute Pre-Routing Hooks @@ -13794,7 +13817,7 @@ class Router: start_time: Final = time.time() if strategy == "simple-shuffle": shuffled: Final = simple_shuffle( - resolve_model_alias=self._get_model_from_alias, + resolve_model_alias=self.get_model_from_alias, healthy_deployments=healthy_deployments, model=model, request_kwargs=request_kwargs, @@ -13892,7 +13915,7 @@ class Router: specific_deployment: bool | None, ): try: - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(request_kwargs) # 1. Execute pre-routing hook responses_call: Final = input is not None and messages is None @@ -13963,7 +13986,7 @@ class Router: start_time: Final = time.perf_counter() if strategy == "simple-shuffle": return simple_shuffle( - resolve_model_alias=self._get_model_from_alias, + resolve_model_alias=self.get_model_from_alias, healthy_deployments=pass_through_deployments, model=model, request_kwargs=request_kwargs, @@ -14126,7 +14149,7 @@ class Router: if not candidates: return None - request_tags: Final = _get_tags_from_request_kwargs(request_kwargs) + request_tags: Final = get_tags_from_request_kwargs(request_kwargs) if request_tags: for tagged in candidates: if tagged.tags and is_valid_deployment_tag( @@ -14223,7 +14246,7 @@ class Router: bound_model: Final = await self._get_claude_code_session_router_binding(cache_key) if not isinstance(bound_model, str): return registered_model_name - bound_registered_model: Final = self._get_model_from_alias(model=bound_model) or bound_model + bound_registered_model: Final = self.get_model_from_alias(model=bound_model) or bound_model if self._select_pre_routing_strategy(bound_registered_model, request_kwargs) is None: await self._delete_claude_code_session_router_binding(cache_key) return registered_model_name @@ -14265,7 +14288,7 @@ class Router: strategy picked. """ self._record_routing_decision(request_kwargs=request_kwargs, routing_decision=None) - requested_registered_model_name: Final = self._get_model_from_alias(model=model) or model + requested_registered_model_name: Final = self.get_model_from_alias(model=model) or model registered_model_name: Final = await self._resolve_claude_code_session_router( model=model, registered_model_name=requested_registered_model_name, @@ -14359,7 +14382,7 @@ class Router: llm_router=self, model_alias=registered_model_name, request_kwargs=request_kwargs, - request_tags=_get_tags_from_request_kwargs(request_kwargs), + request_tags=get_tags_from_request_kwargs(request_kwargs), ) # Shared compression already ran in the pre-call hook, so reuse it rather than # compressing twice. Conditional on arming having actually happened: only the @@ -14408,7 +14431,7 @@ class Router: value=self._consumed_request_tags_stamp( selected_strategy=selected_strategy, pre_routing_hook_response=pre_routing_hook_response, - request_tags=_get_tags_from_request_kwargs(request_kwargs), + request_tags=get_tags_from_request_kwargs(request_kwargs), ), ) @@ -14672,7 +14695,7 @@ class Router: phase_event(_DEPLOYMENT_SELECTED_EVENT, pick_attributes) return healthy_deployments - parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(request_kwargs) # Health-check-based filtering (before cooldown) healthy_deployments = self._filter_health_check_unhealthy_deployments( @@ -14680,7 +14703,7 @@ class Router: parent_otel_span=parent_otel_span, ) - cooldown_deployments: Final = _get_cooldown_deployments( + cooldown_deployments: Final = get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) _pre_cooldown_deployments: Final = healthy_deployments @@ -14737,7 +14760,7 @@ class Router: _cooldown_time = self.cooldown_cache.get_min_cooldown( model_ids=model_ids, parent_otel_span=parent_otel_span ) - _cooldown_list = _get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) + _cooldown_list = get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) raise RouterRateLimitError( model=model, cooldown_time=_cooldown_time, @@ -14750,7 +14773,7 @@ class Router: # if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm ############## Check 'weight' param set for weighted pick ################# shuffled: Final = simple_shuffle( - resolve_model_alias=self._get_model_from_alias, + resolve_model_alias=self.get_model_from_alias, healthy_deployments=healthy_deployments, model=model, request_kwargs=request_kwargs, @@ -14773,7 +14796,7 @@ class Router: _cooldown_time = self.cooldown_cache.get_min_cooldown( model_ids=model_ids, parent_otel_span=parent_otel_span ) - _cooldown_list = _get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) + _cooldown_list = get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) raise RouterRateLimitError( model=model, cooldown_time=_cooldown_time, @@ -14866,12 +14889,12 @@ class Router: ) # 4. Apply health-check and cooldown filtering - parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(request_kwargs) pass_through_deployments = self._filter_health_check_unhealthy_deployments( healthy_deployments=pass_through_deployments, parent_otel_span=parent_otel_span, ) - cooldown_deployments: Final = _get_cooldown_deployments( + cooldown_deployments: Final = get_cooldown_deployments( litellm_router_instance=self, parent_otel_span=parent_otel_span ) pass_through_deployments = self._filter_cooldown_deployments( @@ -14901,7 +14924,7 @@ class Router: _cooldown_time = self.cooldown_cache.get_min_cooldown( model_ids=model_ids, parent_otel_span=parent_otel_span ) - _cooldown_list = _get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) + _cooldown_list = get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) raise RouterRateLimitError( model=model, cooldown_time=_cooldown_time, @@ -14913,7 +14936,7 @@ class Router: # 6. Apply load balancing strategy if strategy == "simple-shuffle": return simple_shuffle( - resolve_model_alias=self._get_model_from_alias, + resolve_model_alias=self.get_model_from_alias, healthy_deployments=pass_through_deployments, model=model, request_kwargs=request_kwargs, @@ -14936,7 +14959,7 @@ class Router: _cooldown_time = self.cooldown_cache.get_min_cooldown( model_ids=model_ids, parent_otel_span=parent_otel_span ) - _cooldown_list = _get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) + _cooldown_list = get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) raise RouterRateLimitError( model=model, cooldown_time=_cooldown_time, @@ -14989,7 +15012,7 @@ class Router: ] @staticmethod - def _is_deployment_blocked(deployment: "Deployment") -> bool: + def is_deployment_blocked(deployment: "Deployment") -> bool: """ Returns True when a `Deployment` Pydantic instance carries the admin-paused flag. Used by credential-lookup helpers so passthrough file / batch endpoints @@ -15000,6 +15023,8 @@ class Router: return False return getattr(model_info, "blocked", None) is True + _is_deployment_blocked = is_deployment_blocked + async def _async_filter_health_check_unhealthy_deployments( self, healthy_deployments: list[dict], diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index 95c981502e7..f4860c4f097 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -86,6 +86,14 @@ class _FeedbackContext: class AdaptiveRouter: """One instance per router_name. Holds in-memory caches + the update queue.""" + @property + def _state_loaded(self) -> bool: + return self.state_loaded + + @_state_loaded.setter + def _state_loaded(self, value: bool) -> None: + self.state_loaded = value + def __init__( self, router_name: str, @@ -109,7 +117,7 @@ class AdaptiveRouter: self._response_signal_updates_total: int = 0 # Set to True once the proxy flusher has loaded persisted priors from # Postgres. Checked to support lazy-load on hot-reloaded routers. - self._state_loaded: bool = False + self.state_loaded: bool = False self._lock = asyncio.Lock() self._init_cold_start_cells() diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 98601038014..8ba1da824c6 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -37,9 +37,9 @@ from litellm.litellm_core_utils.core_helpers import ( get_metadata_variable_name_from_kwargs, ) from litellm.litellm_core_utils.duration_parser import duration_in_seconds -from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs +from litellm.router_strategy.tag_based_routing import get_tags_from_request_kwargs from litellm.router_utils.cooldown_callbacks import ( - _get_prometheus_logger_from_callbacks, + get_prometheus_logger_from_callbacks, ) from litellm.types.llms.openai import AllMessageValues from litellm.types.router import DeploymentTypedDict, LiteLLM_Params, RouterErrors @@ -190,7 +190,7 @@ class RouterBudgetLimiting(CustomLogger): deployment_providers=deployment_providers, spend_map=spend_map, potential_deployments=potential_deployments, - request_tags=_get_tags_from_request_kwargs( + request_tags=get_tags_from_request_kwargs( request_kwargs=request_kwargs, metadata_variable_name=get_metadata_variable_name_from_kwargs(request_kwargs or {}), ), @@ -317,7 +317,7 @@ class RouterBudgetLimiting(CustomLogger): # Resolve tags once before the loop (loop-invariant) _request_tags: list[str] = [] if self.tag_budget_config: - _request_tags = _get_tags_from_request_kwargs( + _request_tags = get_tags_from_request_kwargs( request_kwargs=request_kwargs, metadata_variable_name=get_metadata_variable_name_from_kwargs(request_kwargs or {}), ) @@ -532,7 +532,7 @@ class RouterBudgetLimiting(CustomLogger): response_cost=response_cost, ) - request_tags: Final = _get_tags_from_request_kwargs( + request_tags: Final = get_tags_from_request_kwargs( kwargs, metadata_variable_name=get_metadata_variable_name_from_kwargs(kwargs or {}), ) @@ -747,7 +747,7 @@ class RouterBudgetLimiting(CustomLogger): This is helpful for debugging and monitoring provider budget limits. """ - prometheus_logger: Final = _get_prometheus_logger_from_callbacks() + prometheus_logger: Final = get_prometheus_logger_from_callbacks() if prometheus_logger: prometheus_logger.track_provider_remaining_budget( provider=provider, diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index d7b88476f62..4177771987c 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -45,8 +45,8 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.otel.runtime import phase_event from litellm.litellm_core_utils.classifier_logging import masked_originating_request from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, get_metadata_variable_name_from_kwargs, + get_parent_otel_span_from_kwargs, is_codex_user_agent, ) from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata @@ -3863,7 +3863,7 @@ class ComplexityRouter(CustomLogger): request_kwargs=probe_kwargs, messages=messages, input=input, - parent_otel_span=_get_parent_otel_span_from_kwargs(request_kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(request_kwargs), health_check_probe=True, ) except (RouterRateLimitError, RouterRateLimitErrorBasic, BadRequestError) as exc: diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 622919e3443..83bb22435e6 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -13,7 +13,7 @@ from litellm import ModelResponse, token_counter, verbose_logger from litellm._internal_context import with_service_target from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs, safe_divide_seconds +from litellm.litellm_core_utils.core_helpers import get_parent_otel_span_from_kwargs, safe_divide_seconds from litellm.router_utils.batch_utils import is_batch_retrieve_call_type from litellm.types.utils import LiteLLMPydanticObjectBase @@ -134,7 +134,7 @@ class LowestLatencyLoggingHandler(CustomLogger): # ------------ # Update usage # ------------ - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) request_count_dict: Final = ( self.router_cache.get_cache(key=latency_key, parent_otel_span=parent_otel_span) or {} ) @@ -317,7 +317,7 @@ class LowestLatencyLoggingHandler(CustomLogger): # ------------ # Update usage # ------------ - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) request_count_dict: Final = ( await self.router_cache.async_get_cache( key=latency_key, @@ -513,7 +513,7 @@ class LowestLatencyLoggingHandler(CustomLogger): # get list of potential deployments latency_key: Final = f"{model_group}_map" - parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(request_kwargs) request_count_dict: Final = ( await self.router_cache.async_get_cache(key=latency_key, parent_otel_span=parent_otel_span) or {} ) @@ -542,7 +542,7 @@ class LowestLatencyLoggingHandler(CustomLogger): # get list of potential deployments latency_key: Final = f"{model_group}_map" - parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(request_kwargs) + parent_otel_span: Final[Span | None] = get_parent_otel_span_from_kwargs(request_kwargs) request_count_dict = self.router_cache.get_cache(key=latency_key, parent_otel_span=parent_otel_span) or {} return self._get_available_deployments( diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 9839c9be469..2b37e7dfb40 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -15,7 +15,7 @@ from litellm._internal_context import with_service_target from litellm._logging import verbose_logger, verbose_router_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs +from litellm.litellm_core_utils.core_helpers import get_parent_otel_span_from_kwargs from litellm.router_utils.batch_utils import is_batch_retrieve_call_type from litellm.types.router import RouterErrors from litellm.types.utils import LiteLLMPydanticObjectBase, StandardLoggingPayload @@ -325,7 +325,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): # Update usage # ------------ # update cache - parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) + parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) ## TPM await self.router_cache.async_increment_cache_post_call( key=tpm_key, diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index 50dce250920..541fdd19cff 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -383,7 +383,7 @@ def _all_deployments_or_fallback( fallback: _DeploymentPool, ) -> Sequence[_DeploymentLike | DeploymentTypedDict] | Mapping[_DeploymentLike, object]: try: - return llm_router_instance._get_all_deployments(model_name=model) + return llm_router_instance.get_all_deployments(model_name=model) except Exception: # noqa: BLE001 # fail safe toward today's healthy-only behavior on lookup errors return fallback @@ -442,7 +442,7 @@ def _tag_known_to_group( if tag_set & routing_confirmed: return True try: - all_deployments: Final = llm_router_instance._get_all_deployments(model_name=model) + all_deployments: Final = llm_router_instance.get_all_deployments(model_name=model) except Exception: # noqa: BLE001 # fail safe toward "unrecognized" so lookup errors preserve the existing silent-fallback behavior return False return any( @@ -665,9 +665,9 @@ def _tags_in_metadata(metadata: object, key: str = "tags") -> list[str]: return [tag for tag in typed_tags if isinstance(tag, str)] -def _get_tags_from_request_kwargs( +def get_tags_from_request_kwargs( request_kwargs: Mapping[str, object] | None = None, - metadata_variable_name: Literal["metadata", "litellm_metadata"] | None = None, + metadata_variable_name: str | None = None, ) -> list[str]: """ Helper to get tags from request kwargs @@ -693,3 +693,6 @@ def _get_tags_from_request_kwargs( typed_litellm_params: Final[Mapping[str, object]] = litellm_params return _tags_in_metadata(typed_litellm_params.get(resolved_variable_name)) return [] + + +_get_tags_from_request_kwargs = get_tags_from_request_kwargs diff --git a/litellm/router_utils/add_retry_fallback_headers.py b/litellm/router_utils/add_retry_fallback_headers.py index 6e07693b7ea..ee5d6c1194b 100644 --- a/litellm/router_utils/add_retry_fallback_headers.py +++ b/litellm/router_utils/add_retry_fallback_headers.py @@ -1,8 +1,8 @@ import json import math -from collections.abc import Mapping +from collections.abc import Awaitable, Mapping from types import MappingProxyType -from typing import Any, Final, Protocol, TypedDict, cast +from typing import Final, Protocol, TypedDict, cast from pydantic import BaseModel, TypeAdapter, ValidationError @@ -14,10 +14,17 @@ class FallbackErrorInfo(TypedDict): code: str | None -class _HiddenParamsHost(Protocol): +class HiddenParamsHost(Protocol): _hidden_params: dict[str, object] +_HiddenParamsHost = HiddenParamsHost + + +class AsyncIteratorProtocol(Protocol): + def __anext__(self) -> Awaitable[object]: ... + + _EMPTY_OBJECT_MAPPING: Final[Mapping[str, object]] = MappingProxyType({}) _ROUTING_HEADER_MAPPING: Final = TypeAdapter(Mapping[str, object]) _COMPLEXITY_ROUTER_HEADER_PREFIX: Final = "x-litellm-complexity-router-" @@ -90,17 +97,27 @@ class HiddenParamsAsyncIteratorWrapper: """ def __init__(self, inner: object) -> None: - self._inner = inner + self.inner = inner self._hidden_params: dict[str, object] = {} + @property + def _inner(self) -> object: + return self.inner + + @_inner.setter + def _inner(self, value: object) -> None: + self.inner = value + def __aiter__(self) -> "HiddenParamsAsyncIteratorWrapper": return self async def __anext__(self) -> object: - return await cast(Any, self._inner).__anext__() + return await cast( # cast-ok: provider stream is guarded by __anext__ + AsyncIteratorProtocol, self.inner + ).__anext__() async def aclose(self) -> None: - aclose: Final = getattr(self._inner, "aclose", None) + aclose: Final = getattr(self.inner, "aclose", None) if callable(aclose): await aclose() @@ -213,7 +230,7 @@ def _write_hidden_params(response: object, hidden_params: dict[str, object]) -> if isinstance(response, dict): response["_hidden_params"] = hidden_params elif hasattr(response, "_hidden_params"): - cast(_HiddenParamsHost, response)._hidden_params = hidden_params + setattr(response, "_hidden_params", hidden_params) def _ensure_additional_headers_dict( diff --git a/litellm/router_utils/batch_utils.py b/litellm/router_utils/batch_utils.py index 386a2135239..d0725addb30 100644 --- a/litellm/router_utils/batch_utils.py +++ b/litellm/router_utils/batch_utils.py @@ -150,7 +150,7 @@ def replace_model_in_jsonl(file_content: FileTypes, new_model_name: str) -> File return file_content -def _get_router_metadata_variable_name(function_name: str | None) -> str: +def get_router_metadata_variable_name(function_name: str | None) -> str: """ Helper to return what the "metadata" field should be called in the request data @@ -173,6 +173,9 @@ def _get_router_metadata_variable_name(function_name: str | None) -> str: return "metadata" +_get_router_metadata_variable_name = get_router_metadata_variable_name + + BATCH_RETRIEVE_CALL_TYPES: Final = frozenset( { CallTypes.aretrieve_batch.value, diff --git a/litellm/router_utils/common_utils.py b/litellm/router_utils/common_utils.py index 7d111a80264..6e400738b8e 100644 --- a/litellm/router_utils/common_utils.py +++ b/litellm/router_utils/common_utils.py @@ -17,7 +17,7 @@ from litellm.types.router import CredentialLiteLLMParams from litellm.types.utils import LlmProviders -def _is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool: +def is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool: if request_kwargs is None: return False metadata_value: Final = request_kwargs.get("metadata") @@ -28,6 +28,9 @@ def _is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool return getattr(user_api_key_auth, "user_role", None) == "proxy_admin" +_is_proxy_admin_request = is_proxy_admin_request + + def get_request_team_id(request_kwargs: Mapping[str, object] | None) -> str | None: """The caller's team id, from whichever metadata bucket this surface writes to.""" if request_kwargs is None: @@ -155,7 +158,7 @@ def filter_team_based_models( metadata: Final = request_kwargs.get("metadata") or {} litellm_metadata: Final = request_kwargs.get("litellm_metadata") or {} request_team_id: Final = get_request_team_id(request_kwargs) - if request_team_id is None and _is_proxy_admin_request(request_kwargs) and isinstance(healthy_deployments, list): + if request_team_id is None and is_proxy_admin_request(request_kwargs) and isinstance(healthy_deployments, list): requested_model: Final = ( request_kwargs.get("model") or metadata.get("model_group") or litellm_metadata.get("model_group") ) diff --git a/litellm/router_utils/cooldown_cache.py b/litellm/router_utils/cooldown_cache.py index f780bb3364a..aeab916a8cb 100644 --- a/litellm/router_utils/cooldown_cache.py +++ b/litellm/router_utils/cooldown_cache.py @@ -84,7 +84,7 @@ class CooldownCache: # Store the cooldown information for the deployment separately cooldown_data: Final = CooldownCacheValue( - exception_received=self.exception_masker._mask_value(str(original_exception)), + exception_received=self.exception_masker.mask_value(str(original_exception)), status_code=str(exception_status), timestamp=current_time, cooldown_time=cooldown_time, diff --git a/litellm/router_utils/cooldown_callbacks.py b/litellm/router_utils/cooldown_callbacks.py index 9773964461b..65aa01750ac 100644 --- a/litellm/router_utils/cooldown_callbacks.py +++ b/litellm/router_utils/cooldown_callbacks.py @@ -57,7 +57,7 @@ async def router_cooldown_event_callback( pass # get the prometheus logger from in memory loggers - prometheusLogger: Final[PrometheusLogger | None] = _get_prometheus_logger_from_callbacks() + prometheusLogger: Final[PrometheusLogger | None] = get_prometheus_logger_from_callbacks() if prometheusLogger is not None: prometheusLogger.set_deployment_complete_outage( @@ -78,7 +78,7 @@ async def router_cooldown_event_callback( return -def _get_prometheus_logger_from_callbacks() -> PrometheusLogger | None: +def get_prometheus_logger_from_callbacks() -> PrometheusLogger | None: """ Checks if prometheus is a initalized callback, if yes returns it """ @@ -95,3 +95,6 @@ def _get_prometheus_logger_from_callbacks() -> PrometheusLogger | None: return global_callback return None + + +_get_prometheus_logger_from_callbacks = get_prometheus_logger_from_callbacks diff --git a/litellm/router_utils/cooldown_handlers.py b/litellm/router_utils/cooldown_handlers.py index fcafdfb7402..d9ea1bc8750 100644 --- a/litellm/router_utils/cooldown_handlers.py +++ b/litellm/router_utils/cooldown_handlers.py @@ -24,7 +24,7 @@ from litellm.constants import ( INTERNAL_CALL_ORIGIN_METADATA_KEY, SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD, ) -from litellm.router_utils.cooldown_cache import ROUTER_COOLDOWNS_TARGET +from litellm.router_utils.cooldown_cache import ROUTER_COOLDOWNS_TARGET, CooldownCacheValue from litellm.router_utils.cooldown_callbacks import router_cooldown_event_callback from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN @@ -412,7 +412,7 @@ def _should_cooldown_deployment( # Only apply error rate cooldown when we have enough requests to make the percentage meaningful return True - elif litellm._should_retry(status_code=cast_exception_status_to_int(exception_status)) is False: + elif litellm.should_retry(status_code=cast_exception_status_to_int(exception_status)) is False: return True return False @@ -426,7 +426,7 @@ def _should_cooldown_deployment( return False -def _set_cooldown_deployments( +def set_cooldown_deployments( litellm_router_instance: LitellmRouter, original_exception: Exception, exception_status: str | int, @@ -491,12 +491,15 @@ def _set_cooldown_deployments( return False -async def _async_get_cooldown_deployments( +_set_cooldown_deployments = set_cooldown_deployments + + +async def async_get_cooldown_deployments( litellm_router_instance: LitellmRouter, parent_otel_span: Span | None, ) -> list[str]: """ - Async implementation of '_get_cooldown_deployments' + Async implementation of 'get_cooldown_deployments' """ model_ids: Final = litellm_router_instance.get_model_ids() cooldown_models: Final = await litellm_router_instance.cooldown_cache.async_get_active_cooldowns( @@ -517,12 +520,15 @@ async def _async_get_cooldown_deployments( return cached_value_deployment_ids -async def _async_get_cooldown_deployments_with_debug_info( +_async_get_cooldown_deployments = async_get_cooldown_deployments + + +async def async_get_cooldown_deployments_with_debug_info( litellm_router_instance: LitellmRouter, parent_otel_span: Span | None, -) -> list[tuple]: +) -> list[tuple[str, CooldownCacheValue]]: """ - Async implementation of '_get_cooldown_deployments' + Async implementation of 'get_cooldown_deployments' """ model_ids: Final = litellm_router_instance.get_model_ids() cooldown_models: Final = await litellm_router_instance.cooldown_cache.async_get_active_cooldowns( @@ -533,7 +539,10 @@ async def _async_get_cooldown_deployments_with_debug_info( return cooldown_models -def _get_cooldown_deployments(litellm_router_instance: LitellmRouter, parent_otel_span: Span | None) -> list[str]: +_async_get_cooldown_deployments_with_debug_info = async_get_cooldown_deployments_with_debug_info + + +def get_cooldown_deployments(litellm_router_instance: LitellmRouter, parent_otel_span: Span | None) -> list[str]: """ Get the list of models being cooled down for this minute """ @@ -560,6 +569,9 @@ def _get_cooldown_deployments(litellm_router_instance: LitellmRouter, parent_ote return cached_value_deployment_ids +_get_cooldown_deployments = get_cooldown_deployments + + def should_cooldown_based_on_allowed_fails_policy( litellm_router_instance: LitellmRouter, deployment: str, diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 432351e58f3..e33897e3b63 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -17,13 +17,13 @@ from litellm.router_utils.add_retry_fallback_headers import ( add_fallback_headers_to_response, get_fallback_error_info, ) -from litellm.router_utils.batch_utils import _get_router_metadata_variable_name +from litellm.router_utils.batch_utils import get_router_metadata_variable_name from litellm.router_utils.cooldown_handlers import ( _first_present, # pyright: ignore[reportPrivateUsage] - shared internal helper, used across router_utils - _set_cooldown_deployments, # pyright: ignore[reportPrivateUsage] - shared helper, used across router_utils cast_exception_status_to_int, is_advisor_orchestration_failure, is_caller_timeout_408, + set_cooldown_deployments, ) from litellm.router_utils.router_callbacks.track_deployment_metrics import ( increment_deployment_failures_for_current_minute, @@ -141,7 +141,7 @@ def _trigger_cooldown_for_failed_deployment( litellm_router_instance=litellm_router, deployment_id=deployment_id, ) - _set_cooldown_deployments( + set_cooldown_deployments( litellm_router_instance=litellm_router, exception_status=exception_status, original_exception=exception, @@ -696,7 +696,7 @@ async def run_async_fallback( error_from_fallbacks = original_exception fallback_errors = (get_fallback_error_info(original_exception),) - metadata_variable_name: Final = _get_router_metadata_variable_name( + metadata_variable_name: Final = get_router_metadata_variable_name( function_name=getattr(kwargs.get("original_function"), "__name__", None) ) same_model_group_only: Final = references_provider_scoped_resource(kwargs) or creates_provider_scoped_resource( @@ -855,7 +855,7 @@ async def log_failure_fallback_event(original_model_group: str, kwargs: dict, or verbose_router_logger.error("Error in log_failure_fallback_event: %s", e) -def _check_non_standard_fallback_format(fallbacks: Sequence[object] | None) -> bool: +def check_non_standard_fallback_format(fallbacks: Sequence[object] | None) -> bool: """ Checks if the fallbacks list is a list of strings or a list of dictionaries. @@ -882,5 +882,8 @@ def _check_non_standard_fallback_format(fallbacks: Sequence[object] | None) -> b return False +_check_non_standard_fallback_format = check_non_standard_fallback_format + + def run_non_standard_fallback_format(fallbacks: Sequence[str] | Sequence[Mapping[str, object]], model_group: str): pass diff --git a/litellm/router_utils/handle_error.py b/litellm/router_utils/handle_error.py index bfe02675162..af637f1bdad 100644 --- a/litellm/router_utils/handle_error.py +++ b/litellm/router_utils/handle_error.py @@ -3,7 +3,7 @@ from typing import TYPE_CHECKING, Any, Final from litellm._logging import redact_secrets, verbose_router_logger from litellm.constants import MAX_EXCEPTION_MESSAGE_LENGTH from litellm.router_utils.cooldown_handlers import ( - _async_get_cooldown_deployments_with_debug_info, + async_get_cooldown_deployments_with_debug_info, ) from litellm.types.integrations.slack_alerting import AlertType from litellm.types.router import RouterRateLimitError @@ -79,7 +79,7 @@ async def async_raise_no_deployment_exception( _cooldown_time: Final = litellm_router_instance.cooldown_cache.get_min_cooldown( model_ids=model_ids, parent_otel_span=parent_otel_span ) - _cooldown_list: Final = await _async_get_cooldown_deployments_with_debug_info( + _cooldown_list: Final = await async_get_cooldown_deployments_with_debug_info( litellm_router_instance=litellm_router_instance, parent_otel_span=parent_otel_span, ) diff --git a/litellm/router_utils/pattern_match_deployments.py b/litellm/router_utils/pattern_match_deployments.py index cbb19cec9c2..7ce28a3f606 100644 --- a/litellm/router_utils/pattern_match_deployments.py +++ b/litellm/router_utils/pattern_match_deployments.py @@ -69,7 +69,7 @@ class PatternMatchRouter: llm_deployment: str or List[str] """ # Convert the pattern to a regex - regex: Final = self._pattern_to_regex(pattern) + regex: Final = self.pattern_to_regex(pattern) if regex in self.patterns: self.patterns[regex].append(llm_deployment) return @@ -86,7 +86,7 @@ class PatternMatchRouter: if (remaining := [d for d in deployments if (d.get("model_info") or {}).get("id") != model_id]) } - def _pattern_to_regex(self, pattern: str) -> str: + def pattern_to_regex(self, pattern: str) -> str: """ Convert a wildcard pattern to a regex pattern @@ -110,6 +110,8 @@ class PatternMatchRouter: # return f"^{regex}$" return re.escape(pattern).replace(r"\*", "(.*)") + _pattern_to_regex = pattern_to_regex + def _return_pattern_matched_deployments(self, matched_pattern: Match, deployments: list[dict]) -> list[dict]: new_deployments: Final = [] for deployment in deployments: @@ -141,7 +143,7 @@ class PatternMatchRouter: return None regex_filtered_model_names: Final = ( - tuple(self._pattern_to_regex(m) for m in filtered_model_names) + tuple(self.pattern_to_regex(m) for m in filtered_model_names) if filtered_model_names is not None else () ) diff --git a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py index cf1f18abcba..5d7241db379 100644 --- a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py @@ -131,7 +131,7 @@ class EncryptedContentAffinityCheck(CustomLogger): item_id: Final = item.get("id") if item_id and isinstance(item_id, str): - decoded: Final = ResponsesAPIRequestUtils._decode_encrypted_item_id(item_id) + decoded: Final = ResponsesAPIRequestUtils.decode_encrypted_item_id(item_id) if decoded: return decoded.get("model_id") @@ -156,7 +156,7 @@ class EncryptedContentAffinityCheck(CustomLogger): @staticmethod def _model_id_from_wrapped_encrypted_content(encrypted_content: str) -> str | None: - model_id, _ = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(encrypted_content) + model_id, _ = ResponsesAPIRequestUtils.unwrap_encrypted_content_with_model_id(encrypted_content) return model_id or None @staticmethod diff --git a/litellm/router_utils/pre_call_checks/model_rate_limit_check.py b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py index 3f911bf0825..17e7ec176ab 100644 --- a/litellm/router_utils/pre_call_checks/model_rate_limit_check.py +++ b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py @@ -332,7 +332,7 @@ class ModelRateLimitingCheck(CustomLogger): @with_service_target(ROUTER_USAGE_TARGET) async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, + get_parent_otel_span_from_kwargs, ) try: @@ -350,7 +350,7 @@ class ModelRateLimitingCheck(CustomLogger): self.dual_cache, kwargs, response_obj, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) # Fall through: a deployment can also configure tpm/rpm alongside # itpm/otpm, and that path's pre-call check reads the tpm_key @@ -390,7 +390,7 @@ class ModelRateLimitingCheck(CustomLogger): async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, + get_parent_otel_span_from_kwargs, ) # Never fail the primary logging pipeline over an io-token refund error. @@ -398,7 +398,7 @@ class ModelRateLimitingCheck(CustomLogger): await async_io_token_refund_failure( self.dual_cache, kwargs, - parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs), + parent_otel_span=get_parent_otel_span_from_kwargs(kwargs), ) @with_service_target(ROUTER_USAGE_TARGET) diff --git a/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py b/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py index 8ede0109e92..9c89e9b7ef1 100644 --- a/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py +++ b/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py @@ -43,7 +43,7 @@ class ResponsesApiDeploymentCheck(CustomLogger): if previous_response_id is None: return healthy_deployments - decoded_response: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id( + decoded_response: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id( response_id=previous_response_id, ) model_id: Final = decoded_response.get("model_id") diff --git a/litellm/rust_bridge/callbacks_legacy_python.py b/litellm/rust_bridge/callbacks_legacy_python.py index 46af4e236eb..bfd16b29bf6 100644 --- a/litellm/rust_bridge/callbacks_legacy_python.py +++ b/litellm/rust_bridge/callbacks_legacy_python.py @@ -227,7 +227,7 @@ def post_call( def defers_async_logging(logger: LoggingSurface) -> bool: - return bool(getattr(logger, "_defer_async_logging", False)) + return bool(getattr(logger, "defer_async_logging", False)) def defer_success(logger: LoggingSurface, pending: object) -> None: diff --git a/litellm/types/agents.py b/litellm/types/agents.py index badaf20dba0..b3aa86c249d 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -413,7 +413,7 @@ class MakeAgentsPublicRequest(LiteLLMBaseModel): agent_ids: list[str] -def _normalize_a2a_jsonrpc_response( +def normalize_a2a_jsonrpc_response( response_dict: Mapping[str, object], request_id: object | None = None, ) -> dict[str, object]: @@ -443,6 +443,9 @@ def _normalize_a2a_jsonrpc_response( return normalized +_normalize_a2a_jsonrpc_response = normalize_a2a_jsonrpc_response + + class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase): """ LiteLLM wrapper for A2A SendMessageResponse. @@ -481,7 +484,7 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase): Returns: LiteLLMSendMessageResponse with _hidden_params support """ - response_dict: Final = _normalize_a2a_jsonrpc_response( + response_dict: Final = normalize_a2a_jsonrpc_response( response.model_dump(mode="json", exclude_none=True), request_id=request_id ) return cls.model_validate(response_dict) @@ -502,4 +505,4 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase): Returns: LiteLLMSendMessageResponse with _hidden_params support """ - return cls.model_validate(_normalize_a2a_jsonrpc_response(response_dict, request_id=request_id)) + return cls.model_validate(normalize_a2a_jsonrpc_response(response_dict, request_id=request_id)) diff --git a/litellm/types/completion.py b/litellm/types/completion.py index 156ae96410b..97b77b4ca01 100644 --- a/litellm/types/completion.py +++ b/litellm/types/completion.py @@ -206,7 +206,7 @@ class CompletionRequest(LiteLLMBaseModel): @dataclass(frozen=True, slots=True) -class _CompletionDispatchContext: +class CompletionDispatchContext: _azure_detection_model: str acompletion: bool api_base: str | None @@ -240,6 +240,9 @@ class _CompletionDispatchContext: top_p: float | None +_CompletionDispatchContext = CompletionDispatchContext + + _CompletionDispatchResult = Union[ Coroutine[Any, Any, Union["ModelResponse", "CustomStreamWrapper"]], "ModelResponse", diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 8f4ad26a4fa..dc552a730ea 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -8,7 +8,7 @@ from typing import Any, ClassVar, Final, Literal, cast import litellm -def _sanitize_prometheus_label_name(label: str) -> str: +def sanitize_prometheus_label_name(label: str) -> str: """ Sanitize a label name to comply with Prometheus label name requirements. @@ -40,11 +40,14 @@ def _sanitize_prometheus_label_name(label: str) -> str: return sanitized +_sanitize_prometheus_label_name = sanitize_prometheus_label_name + + # v1: single translate pass + escape loop (avoids chained str.replace allocations). _PROMETHEUS_LABEL_VALUE_TRANSLATE_V1: Final = str.maketrans("\n", " ", "\r\u2028\u2029") -def _sanitize_prometheus_label_value(value: object | None) -> str | None: +def sanitize_prometheus_label_value(value: object | None) -> str | None: """ Same semantics as :func:`_sanitize_prometheus_label_value`, implemented with ``str.translate`` plus a single escape pass instead of chained ``replace``. @@ -70,6 +73,9 @@ def _sanitize_prometheus_label_value(value: object | None) -> str | None: return "".join(parts) +_sanitize_prometheus_label_value = sanitize_prometheus_label_value + + @dataclass class MetricValidationError: """Error for invalid metric name""" @@ -963,11 +969,11 @@ class PrometheusMetricLabels: # Add custom metadata labels custom_labels.extend( - [_sanitize_prometheus_label_name(metric) for metric in litellm.custom_prometheus_metadata_labels] + [sanitize_prometheus_label_name(metric) for metric in litellm.custom_prometheus_metadata_labels] ) # Add custom tags labels - custom_labels.extend([_sanitize_prometheus_label_name(f"tag_{tag}") for tag in litellm.custom_prometheus_tags]) + custom_labels.extend([sanitize_prometheus_label_name(f"tag_{tag}") for tag in litellm.custom_prometheus_tags]) # Conditionally add stream label to litellm_proxy_total_requests_metric if ( diff --git a/litellm/types/utils.py b/litellm/types/utils.py index a0279bcadc7..3ec5e91af14 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -125,10 +125,13 @@ else: VectorStoreSearchResponse = Any -def _generate_id(): # private helper function +def generate_id() -> str: return "chatcmpl-" + str(uuid.uuid4()) +_generate_id = generate_id + + class SafeAttributeModel: """ A base model that provides safe attribute access. @@ -2130,7 +2133,7 @@ class ModelResponseStream(ModelResponseBase): kwargs["choices"] = [StreamingChoices()] if id is None: - id = _generate_id() + id = generate_id() else: id = id if created is None: @@ -2218,7 +2221,7 @@ class ModelResponse(ModelResponseBase): else: choices = [Choices()] if id is None: - id = _generate_id() + id = generate_id() else: id = id if created is None: @@ -2483,7 +2486,7 @@ class TextCompletionResponse(OpenAIObject): if object is not None: object = object if id is None: - id = _generate_id() + id = generate_id() else: id = id if created is None: diff --git a/litellm/utils.py b/litellm/utils.py index f434b8be425..6c0d9064dbe 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -43,9 +43,9 @@ import httpx import openai from httpx import Proxy from httpx._utils import get_environment_proxies -from openai.lib import _parsing, _pydantic +from openai.lib import _parsing, _pydantic # pyright: ignore[reportPrivateUsage] # OpenAI parser module is private from openai.types.chat.completion_create_params import ResponseFormat -from pydantic import BaseModel +from pydantic import BaseModel, JsonValue import litellm import litellm.litellm_core_utils @@ -54,10 +54,10 @@ import litellm.litellm_core_utils import litellm.litellm_core_utils.json_validation_rule from litellm._internal_context import is_internal_call from litellm._lazy_imports import ( - _get_default_encoding, - _get_messages_reach_token_count, _get_modified_max_tokens, - _get_token_counter_new, + get_default_encoding, + get_messages_reach_token_count, + get_token_counter_new, ) from litellm._uuid import uuid from litellm.constants import ( @@ -290,7 +290,7 @@ except (ImportError, AttributeError, TypeError): claude_json_str = json.dumps(json_data) import importlib.metadata from collections.abc import AsyncIterator, Callable, Collection, Iterable, Iterator, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, runtime_checkable +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, cast, runtime_checkable from typing_extensions import assert_never @@ -333,23 +333,23 @@ if TYPE_CHECKING: # These imports allow mypy to understand the types when these are accessed via __getattr__ from litellm.litellm_core_utils.exception_mapping_utils import exception_type from litellm.litellm_core_utils.get_litellm_params import ( - _get_base_model_from_litellm_call_metadata, + get_base_model_from_litellm_call_metadata, get_litellm_params, ) from litellm.litellm_core_utils.get_llm_provider_logic import ( - _is_non_openai_azure_model, get_llm_provider, + is_non_openai_azure_model, ) from litellm.litellm_core_utils.get_supported_openai_params import ( get_supported_openai_params, ) - from litellm.litellm_core_utils.llm_request_utils import _ensure_extra_body_is_safe + from litellm.litellm_core_utils.llm_request_utils import ensure_extra_body_is_safe from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( LiteLLMResponseObjectHandler, - _handle_invalid_parallel_tool_calls, convert_to_model_response_object, convert_to_streaming_response, convert_to_streaming_response_async, + handle_invalid_parallel_tool_calls, ) from litellm.litellm_core_utils.llm_response_utils.get_api_base import get_api_base from litellm.litellm_core_utils.llm_response_utils.get_formatted_prompt import ( @@ -362,9 +362,6 @@ if TYPE_CHECKING: ResponseMetadata, update_response_metadata, ) - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - _parse_content_for_reasoning, - ) from litellm.litellm_core_utils.redact_messages import ( LiteLLMLoggingObject, redact_message_input_output_from_logging, @@ -455,7 +452,7 @@ from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfi from litellm.llms.base_llm.skills.transformation import BaseSkillsAPIConfig from litellm.secret_managers.main import get_secret -from ._logging import _is_debugging_on, verbose_logger +from ._logging import is_debugging_on, verbose_logger from .caching.caching import ( AzureBlobCache, Cache, @@ -580,7 +577,7 @@ def custom_llm_setup(): litellm._custom_providers.append(custom_llm["provider"]) -def _add_custom_logger_callback_to_specific_event(callback: str, logging_event: Literal["success", "failure"]) -> None: +def add_custom_logger_callback_to_specific_event(callback: str, logging_event: Literal["success", "failure"]) -> None: """ Add a custom logger callback to the specific event """ @@ -620,6 +617,9 @@ def _add_custom_logger_callback_to_specific_event(callback: str, logging_event: litellm._async_failure_callback.remove(callback) # remove the string from the callback list +_add_custom_logger_callback_to_specific_event = add_custom_logger_callback_to_specific_event + + def _custom_logger_class_exists_in_success_callbacks( callback_class: CustomLogger, ) -> bool: @@ -719,8 +719,8 @@ def load_credentials_from_list(kwargs: dict): def get_dynamic_callbacks( - dynamic_callbacks: list[str | Callable | CustomLogger] | None, -) -> list: + dynamic_callbacks: list[str | Callable[..., object] | CustomLogger] | None, +) -> list[str | Callable[..., object] | CustomLogger]: returned_callbacks: Final = litellm.callbacks.copy() if dynamic_callbacks: returned_callbacks.extend(dynamic_callbacks) @@ -1059,7 +1059,7 @@ def function_setup( litellm.logging_callback_manager.add_litellm_async_success_callback(callback) removed_async_items.append(index) elif callback in litellm._known_custom_logger_compatible_callbacks and isinstance(callback, str): - _add_custom_logger_callback_to_specific_event(callback, "success") + add_custom_logger_callback_to_specific_event(callback, "success") # Pop the async items from success_callback in reverse order to avoid index issues for index in reversed(removed_async_items): @@ -1072,7 +1072,7 @@ def function_setup( litellm.logging_callback_manager.add_litellm_async_failure_callback(callback) removed_async_items.append(index) elif callback in litellm._known_custom_logger_compatible_callbacks and isinstance(callback, str): - _add_custom_logger_callback_to_specific_event(callback, "failure") + add_custom_logger_callback_to_specific_event(callback, "failure") # Pop the async items from failure_callback in reverse order to avoid index issues for index in reversed(removed_async_items): @@ -1393,12 +1393,12 @@ def _schedule_async_success_logging( ) ) - if not getattr(logging_obj, "_defer_async_logging", False): + if not getattr(logging_obj, "defer_async_logging", False): _enqueue_async_logging() return - if getattr(logging_obj, "_enqueue_deferred_logging", None) is not None: + if getattr(logging_obj, "enqueue_deferred_logging", None) is not None: return - logging_obj._enqueue_deferred_logging = _enqueue_async_logging + logging_obj.enqueue_deferred_logging = _enqueue_async_logging async def _client_async_logging_helper( @@ -1478,9 +1478,9 @@ async def async_pre_call_deployment_hook(kwargs: dict[str, Any], call_type: str) modified_kwargs = kwargs.copy() - CustomLogger: Final = _get_cached_custom_logger() + custom_logger_class: Final = _get_cached_custom_logger() for callback in litellm.callbacks: - if isinstance(callback, CustomLogger): + if isinstance(callback, custom_logger_class): result = await callback.async_pre_call_deployment_hook(modified_kwargs, typed_call_type) if result is not None: modified_kwargs = result @@ -1501,10 +1501,10 @@ async def async_post_call_success_deployment_hook( modified_response = response - CustomLogger: Final = _get_cached_custom_logger() + custom_logger_class: Final = _get_cached_custom_logger() CustomGuardrail: Final = _get_cached_custom_guardrail() for callback in litellm.callbacks: - if isinstance(callback, CustomLogger): + if isinstance(callback, custom_logger_class): try: result = await callback.async_post_call_success_deployment_hook( request_data, cast(LLMResponseTypes, modified_response), typed_call_type @@ -1557,9 +1557,9 @@ async def async_post_call_failure_deployment_hook( safe_request_data: Final = MappingProxyType({k: v for k, v in request_data.items() if k != "attempted_targets"}) safe_exception: Final = _snapshot_exception_for_hook(exception) - CustomLogger: Final = _get_cached_custom_logger() + custom_logger_class: Final = _get_cached_custom_logger() for callback in litellm.callbacks: - if isinstance(callback, CustomLogger): + if isinstance(callback, custom_logger_class): try: if _accepts_fallback_depth_kwarg_for_class(type(callback)): await callback.async_post_call_failure_deployment_hook( @@ -1627,7 +1627,7 @@ def post_call_processing( and optional_params["response_format"].get("json_schema") is not None ): json_response_format = optional_params["response_format"] - elif _parsing._completions.is_basemodel_type( + elif _parsing._completions.is_basemodel_type( # pyright: ignore[reportPrivateUsage] # OpenAI parser helper is private optional_params["response_format"] ): json_response_format = type_to_response_format_param( @@ -1664,6 +1664,16 @@ def post_call_processing( raise e +def _get_call_type_from_public_name(function_name: str) -> str: + match function_name: + case "arealtime": + return CallTypes.arealtime.value + case "aresponses_websocket": + return CallTypes.aresponses_websocket.value + case _: + return function_name + + def _is_litellm_router_call(kwargs: Mapping[str, object], *, is_async: bool) -> bool: """Router completion uses metadata. Async generic calls retry with litellm_metadata; sync generic calls need SDK retries.""" metadata_buckets: Final = ( @@ -1682,7 +1692,7 @@ def client(original_function): def wrapper(*args, **kwargs): # DO NOT MOVE THIS. It always needs to run first # Check if this is an async function. If so only execute the async function - call_type = original_function.__name__ + call_type: Final = _get_call_type_from_public_name(original_function.__name__) if _is_async_request(kwargs): # [OPTIONAL] CHECK MAX RETRIES / REQUEST if max_retries_per_request_hit(kwargs, litellm.num_retries_per_request): @@ -1719,7 +1729,7 @@ def client(original_function): try: if logging_obj is None: logging_obj, kwargs = function_setup( - original_function.__name__, rules_obj, start_time, *args, is_async_call=False, **kwargs + call_type, rules_obj, start_time, *args, is_async_call=False, **kwargs ) # Type assertion: logging_obj is guaranteed to be non-None after function_setup @@ -1734,7 +1744,7 @@ def client(original_function): request_kwargs=kwargs, start_time=start_time, ) - logging_obj._llm_caching_handler = _llm_caching_handler + logging_obj.llm_caching_handler = _llm_caching_handler # [OPTIONAL] CHECK BUDGET if litellm.max_budget: @@ -1771,7 +1781,7 @@ def client(original_function): ): # allow users to control returning cached responses from the completion function # checking cache verbose_logger.debug("INSIDE CHECKING SYNC CACHE") - caching_handler_response: Final[CachingHandlerResponse] = _llm_caching_handler._sync_get_cache( + caching_handler_response: Final[CachingHandlerResponse] = _llm_caching_handler.sync_get_cache( model=model or "", original_function=original_function, logging_obj=logging_obj, @@ -1898,7 +1908,6 @@ def client(original_function): # RETURN RESULT return result except Exception as e: - call_type = original_function.__name__ if call_type == CallTypes.completion.value: num_retries = kwargs.get("num_retries", None) or litellm.num_retries or None if kwargs.get("retry_policy", None): @@ -1973,6 +1982,7 @@ def client(original_function): async def wrapper_async(*args, **kwargs): print_args_passed_to_litellm(original_function, args, kwargs) start_time: Final = datetime.datetime.now() + call_type: Final = _get_call_type_from_public_name(original_function.__name__) result = None _update_response_metadata: Final[_ResponseMetadataUpdater] = litellm_utils.update_response_metadata logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) @@ -1983,7 +1993,6 @@ def client(original_function): start_time=start_time, ) # only set litellm_call_id if its not in kwargs - call_type = original_function.__name__ if "litellm_call_id" not in kwargs: kwargs["litellm_call_id"] = str(uuid.uuid4()) @@ -1995,7 +2004,7 @@ def client(original_function): try: if logging_obj is None: - logging_obj, kwargs = function_setup(original_function.__name__, rules_obj, start_time, *args, **kwargs) + logging_obj, kwargs = function_setup(call_type, rules_obj, start_time, *args, **kwargs) # Type assertion: logging_obj is guaranteed to be non-None after function_setup assert logging_obj is not None, "logging_obj should not be None after function_setup" @@ -2015,7 +2024,7 @@ def client(original_function): kwargs["litellm_logging_obj"] = logging_obj ## LOAD CREDENTIALS load_credentials_from_list(kwargs) - logging_obj._llm_caching_handler = _llm_caching_handler + logging_obj.llm_caching_handler = _llm_caching_handler # [OPTIONAL] CHECK BUDGET if litellm.max_budget: if litellm._current_cost > litellm.max_budget: @@ -2025,11 +2034,11 @@ def client(original_function): ) # [OPTIONAL] CHECK CACHE - if _is_debugging_on(): + if is_debugging_on(): print_verbose( f"ASYNC kwargs[caching]: {kwargs.get('caching', False)}; litellm.cache: {litellm.cache}; kwargs.get('cache'): {kwargs.get('cache', None)}" ) - _caching_handler_response: CachingHandlerResponse | None = await _llm_caching_handler._async_get_cache( + _caching_handler_response: CachingHandlerResponse | None = await _llm_caching_handler.async_get_cache( model=model or "", original_function=original_function, logging_obj=logging_obj, @@ -2177,7 +2186,7 @@ def client(original_function): is_completion_with_fallbacks=is_completion_with_fallbacks, is_litellm_internal_call=_is_litellm_internal_call, ) - return _llm_caching_handler._combine_cached_embedding_response_with_api_result( + return _llm_caching_handler.combine_cached_embedding_response_with_api_result( _caching_handler_response=_caching_handler_response, embedding_response=result, start_time=start_time, @@ -2220,9 +2229,9 @@ def client(original_function): except Exception as e: raise e - call_type = original_function.__name__ + retry_call_type: Final = _get_call_type_from_public_name(original_function.__name__) num_retries, kwargs = _get_wrapper_num_retries(kwargs=kwargs, exception=e) - if call_type == CallTypes.acompletion.value: + if retry_call_type == CallTypes.acompletion.value: context_window_fallback_dict: Final = kwargs.get("context_window_fallback_dict", {}) is_acompletion_litellm_router_call: Final = _is_litellm_router_call(kwargs, is_async=True) @@ -2255,7 +2264,7 @@ def client(original_function): kwargs["model"] = context_window_fallback_dict[model] result = await original_function(*args, **kwargs) return result - elif call_type == CallTypes.aresponses.value: + elif retry_call_type == CallTypes.aresponses.value: is_aresponses_litellm_router_call: Final = _is_litellm_router_call(kwargs, is_async=True) if ( @@ -2366,7 +2375,10 @@ def _is_streaming_request( return call_type in _STREAMING_CALL_TYPES -def _select_tokenizer(model: str, custom_tokenizer: CustomHuggingfaceTokenizer | None = None): +_SchemaT = TypeVar("_SchemaT", bound=JsonValue) + + +def select_tokenizer(model: str, custom_tokenizer: CustomHuggingfaceTokenizer | None = None) -> SelectTokenizerResponse: if custom_tokenizer is not None: return _select_custom_tokenizer_helper( identifier=custom_tokenizer["identifier"], @@ -2377,6 +2389,9 @@ def _select_tokenizer(model: str, custom_tokenizer: CustomHuggingfaceTokenizer | return _select_tokenizer_helper(model=model) +_select_tokenizer = select_tokenizer + + def _huggingface_tokenizer_backend() -> Decision: """The backend `tokenizer_dispatch.from_str` / `from_pretrained` will select right now. @@ -2413,7 +2428,7 @@ def _select_tokenizer_helper(model: str) -> SelectTokenizerResponse: def _return_openai_tokenizer(model: str) -> SelectTokenizerResponse: - return {"type": "openai_tokenizer", "tokenizer": _get_default_encoding()} + return {"type": "openai_tokenizer", "tokenizer": get_default_encoding()} def uses_anthropic_tokenizer(model: str) -> bool: @@ -2474,7 +2489,7 @@ def encode(model="", text="", custom_tokenizer: dict | None = None): Returns: enc: The encoded text. """ - tokenizer_json: Final = custom_tokenizer or _select_tokenizer(model=model) + tokenizer_json: Final = custom_tokenizer or select_tokenizer(model=model) if tokenizer_json["type"] == "openai_tokenizer": openai_tokenizer: Final = cast( # cast-ok: [LIT006] caller's explicit type tag selects this interface Encoding, tokenizer_json["tokenizer"] @@ -2498,7 +2513,7 @@ def decode( LiteLLM round-trip behavior by omitting special tokens by default. Set to False to inspect decoded BOS/EOS tokens. """ - tokenizer_json: Final = custom_tokenizer or _select_tokenizer(model=model) + tokenizer_json: Final = custom_tokenizer or select_tokenizer(model=model) if tokenizer_json["type"] == "huggingface_tokenizer": ids: Final = strip_special_tokens(tokenizer_json["tokenizer"], tokens) if skip_special_tokens else tokens hf_tokenizer: Final = cast( # cast-ok: [LIT006] caller's explicit type tag selects this interface @@ -2566,7 +2581,7 @@ def token_counter( if litellm.disable_token_counter is True: return 0 - return _get_token_counter_new()( + return get_token_counter_new()( model, custom_tokenizer, text, @@ -2605,7 +2620,7 @@ def supports_system_messages(model: str, custom_llm_provider: str | None) -> boo Raises: Exception: If the given model is not found in model_prices_and_context_window.json. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_system_messages", @@ -2626,7 +2641,7 @@ def supports_web_search(model: str, custom_llm_provider: str | None = None) -> b Raises: Exception: If the given model is not found in model_prices_and_context_window.json. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_web_search", @@ -2647,7 +2662,7 @@ def supports_url_context(model: str, custom_llm_provider: str | None = None) -> Raises: Exception: If the given model is not found in model_prices_and_context_window.json. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_url_context", @@ -2673,7 +2688,7 @@ def supports_native_streaming(model: str, custom_llm_provider: str | None) -> bo model=model, custom_llm_provider=custom_llm_provider ) - model_info: Final = _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) + model_info: Final = get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) supports_native_streaming = model_info.get("supports_native_streaming", True) if supports_native_streaming is None: supports_native_streaming = True @@ -2725,7 +2740,7 @@ def supports_response_schema(model: str, custom_llm_provider: str | None = None) if custom_llm_provider in PROVIDERS_GLOBALLY_SUPPORT_RESPONSE_SCHEMA: return True - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_response_schema", @@ -2736,7 +2751,7 @@ def supports_parallel_function_calling(model: str, custom_llm_provider: str | No """ Check if the given model supports parallel tool calls and return a boolean value. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_parallel_function_calling", @@ -2757,7 +2772,7 @@ def supports_function_calling(model: str, custom_llm_provider: str | None = None Raises: Exception: If the given model is not found or there's an error in retrieval. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_function_calling", @@ -2768,7 +2783,7 @@ def supports_tool_choice(model: str, custom_llm_provider: str | None = None) -> """ Check if the given model supports `tool_choice` and return a boolean value. """ - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_tool_choice") + return supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_tool_choice") def _supports_provider_info_factory(model: str, custom_llm_provider: str | None, key: str) -> Literal[True] | None: @@ -2783,7 +2798,7 @@ def _supports_provider_info_factory(model: str, custom_llm_provider: str | None, return None -def _supports_factory(model: str, custom_llm_provider: str | None, key: str) -> bool: +def supports_factory(model: str, custom_llm_provider: str | None, key: str) -> bool: """ Check if the given model supports function calling and return a boolean value. @@ -2809,7 +2824,7 @@ def _supports_factory(model: str, custom_llm_provider: str | None, key: str) -> model=model, custom_llm_provider=custom_llm_provider ) - model_info: Final = _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) + model_info: Final = get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) if model_info.get(key, False) is True: return True @@ -2818,7 +2833,7 @@ def _supports_factory(model: str, custom_llm_provider: str | None, key: str) -> # "deepseek/deepseek-chat") exists but is missing a capability # field, check the bare model-name entry (e.g. "deepseek-chat") # which may carry the complete metadata. See #20885. - bare_model_key: Final = _get_model_cost_key(model) + bare_model_key: Final = get_model_cost_key(model) if bare_model_key is not None: bare_entry: Final = litellm.model_cost.get(bare_model_key) or {} if bare_entry.get(key, False) is True: @@ -2845,6 +2860,9 @@ def _supports_factory(model: str, custom_llm_provider: str | None, key: str) -> return False +_supports_factory = supports_factory + + def declared_value_factory(model: str, custom_llm_provider: str | None, key: str) -> str | None: """Return a string value the model map declares for *key*, or ``None`` when it says nothing. @@ -2863,11 +2881,11 @@ def declared_value_factory(model: str, custom_llm_provider: str | None, key: str resolved: Final = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider) resolved_model: Final = resolved[0] resolved_provider: Final = resolved[1] - model_info: Final = _get_model_info_helper(model=resolved_model, custom_llm_provider=resolved_provider) + model_info: Final = get_model_info_helper(model=resolved_model, custom_llm_provider=resolved_provider) declared: Final = model_info.get(key) if isinstance(declared, str): return declared - bare_model_key: Final = _get_model_cost_key(resolved_model) + bare_model_key: Final = get_model_cost_key(resolved_model) bare_entry: Final = litellm.model_cost.get(bare_model_key) if bare_model_key is not None else None if isinstance(bare_entry, dict): bare_declared: Final = bare_entry.get(key) @@ -2909,12 +2927,12 @@ def is_explicitly_disabled_factory(model: str, custom_llm_provider: str | None, model, custom_llm_provider, _, _ = litellm.get_llm_provider( model=model, custom_llm_provider=custom_llm_provider ) - model_info: Final = _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) + model_info: Final = get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) val: Final = model_info.get(key) if val is False: return True if val is None: - bare_model_key: Final = _get_model_cost_key(model) + bare_model_key: Final = get_model_cost_key(model) if bare_model_key is not None: bare_entry: Final = litellm.model_cost.get(bare_model_key) or {} if bare_entry.get(key) is False: @@ -2933,17 +2951,17 @@ def is_explicitly_disabled_factory(model: str, custom_llm_provider: str | None, def supports_audio_input(model: str, custom_llm_provider: str | None = None) -> bool: """Check if a given model supports audio input in a chat completion call""" - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_input") + return supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_input") def supports_pdf_input(model: str, custom_llm_provider: str | None = None) -> bool: """Check if a given model supports pdf input in a chat completion call""" - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_pdf_input") + return supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_pdf_input") def supports_audio_output(model: str, custom_llm_provider: str | None = None) -> bool: """Check if a given model supports audio output in a chat completion call""" - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_input") + return supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_input") def supports_prompt_caching(model: str, custom_llm_provider: str | None = None) -> bool: @@ -2960,7 +2978,7 @@ def supports_prompt_caching(model: str, custom_llm_provider: str | None = None) Raises: Exception: If the given model is not found or there's an error in retrieval. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_prompt_caching", @@ -2968,7 +2986,7 @@ def supports_prompt_caching(model: str, custom_llm_provider: str | None = None) def supports_prompt_cache_breakpoint(model: str, custom_llm_provider: str | None = None) -> bool: - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_prompt_cache_breakpoint", @@ -2976,7 +2994,7 @@ def supports_prompt_cache_breakpoint(model: str, custom_llm_provider: str | None def supports_thinking_cache_preservation(model: str, custom_llm_provider: str | None = None) -> bool: - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_thinking_cache_preservation", @@ -2997,7 +3015,7 @@ def supports_computer_use(model: str, custom_llm_provider: str | None = None) -> Raises: Exception: If the given model is not found or there's an error in retrieval. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_computer_use", @@ -3024,7 +3042,7 @@ def supports_vision(model: str, custom_llm_provider: str | None = None) -> bool: Returns: bool: True if the model supports vision, False otherwise. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_vision", @@ -3035,11 +3053,11 @@ def supports_reasoning(model: str, custom_llm_provider: str | None = None) -> bo """ Check if the given model supports reasoning and return a boolean value. """ - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_reasoning") + return supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_reasoning") def supports_anthropic_thinking_payload(model: str, custom_llm_provider: str | None = None) -> bool: - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_anthropic_thinking_payload" ) @@ -3048,14 +3066,14 @@ def supports_none_reasoning_effort(model: str, custom_llm_provider: str | None = """ Check if the given model accepts reasoning effort "none" and return a boolean value. """ - return _supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_none_reasoning_effort") + return supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_none_reasoning_effort") def supports_mid_conversation_system(model: str, custom_llm_provider: str | None = None) -> bool: """ Check if the given model accepts a system role message after the leading system block and return a boolean value. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_mid_conversation_system" ) @@ -3064,7 +3082,7 @@ def supports_native_structured_output(model: str, custom_llm_provider: str | Non """ Check if the given model supports native structured outputs and return a boolean value. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_native_structured_output", @@ -3084,7 +3102,7 @@ def get_supported_regions(model: str, custom_llm_provider: str | None = None) -> model=model, custom_llm_provider=custom_llm_provider ) - model_info: Final = _get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) + model_info: Final = get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider) # Get the key used in model_cost to look up supported_regions # since ModelInfoBase doesn't include this field @@ -3118,7 +3136,7 @@ def supports_embedding_image_input(model: str, custom_llm_provider: str | None = """ Check if the given model supports embedding image input and return a boolean value. """ - return _supports_factory( + return supports_factory( model=model, custom_llm_provider=custom_llm_provider, key="supports_embedding_image_input", @@ -3126,24 +3144,29 @@ def supports_embedding_image_input(model: str, custom_llm_provider: str | None = ####### HELPER FUNCTIONS ################ -def _update_dictionary(existing_dict: dict, new_dict: dict) -> dict: +def update_dictionary(existing_dict: dict[str, object], new_dict: Mapping[str, object]) -> dict[str, object]: for k, v in new_dict.items(): if v is not None: # Convert stringified numbers to appropriate numeric types if isinstance(v, str): existing_dict[k] = _convert_stringified_numbers(v) elif isinstance(v, dict): - existing_nested_dict = existing_dict.get(k) - if isinstance(existing_nested_dict, dict): - existing_dict[k] = {**existing_nested_dict, **v} + nested_dict = cast(dict[str, object], v) + existing_nested_value = existing_dict.get(k) + if isinstance(existing_nested_value, dict): + existing_nested_dict = cast(dict[str, object], existing_nested_value) + existing_dict[k] = {**existing_nested_dict, **nested_dict} else: - existing_dict[k] = dict(v) + existing_dict[k] = dict(nested_dict) else: existing_dict[k] = v return existing_dict +_update_dictionary = update_dictionary + + def _convert_stringified_numbers(value): """Convert stringified numbers (including scientific notation) to appropriate numeric types.""" if isinstance(value, str): @@ -3412,7 +3435,7 @@ def register_model( if _cost_field not in _raw_entry and _cost_field not in value: existing_model.pop(_cost_field, None) ## override / add new keys to the existing model cost dictionary - updated_dictionary = _update_dictionary(existing_model, value) + updated_dictionary = update_dictionary(existing_model, value) litellm.model_cost.setdefault(model_cost_key, {}).update(updated_dictionary) # Invalidate case-insensitive lookup map since model_cost was modified @@ -4094,7 +4117,7 @@ def get_optional_params_embeddings( return final_params -def _remove_additional_properties(schema): +def remove_additional_properties(schema: _SchemaT) -> _SchemaT: """ clean out 'additionalProperties = False'. Causes vertexai/gemini OpenAI API Schema errors - https://github.com/langchain-ai/langchainjs/issues/5240 @@ -4107,17 +4130,20 @@ def _remove_additional_properties(schema): # Recursively process all dictionary values for key, value in schema.items(): - _remove_additional_properties(value) + remove_additional_properties(value) elif isinstance(schema, list): # Recursively process all items in the list for item in schema: - _remove_additional_properties(item) + remove_additional_properties(item) return schema -def _remove_strict_from_schema(schema): +_remove_additional_properties = remove_additional_properties + + +def remove_strict_from_schema(schema: _SchemaT) -> _SchemaT: """ Relevant Issues: https://github.com/BerriAI/litellm/issues/6136, https://github.com/BerriAI/litellm/issues/6088 """ @@ -4128,17 +4154,20 @@ def _remove_strict_from_schema(schema): # Recursively process all dictionary values for key, value in schema.items(): - _remove_strict_from_schema(value) + remove_strict_from_schema(value) elif isinstance(schema, list): # Recursively process all items in the list for item in schema: - _remove_strict_from_schema(item) + remove_strict_from_schema(item) return schema -def _remove_json_schema_refs(schema, max_depth=10): +_remove_strict_from_schema = remove_strict_from_schema + + +def remove_json_schema_refs(schema: _SchemaT, max_depth: int = 10) -> _SchemaT: """ Remove JSON schema reference fields like '$id' and '$schema' that can cause issues with some providers. @@ -4161,16 +4190,19 @@ def _remove_json_schema_refs(schema, max_depth=10): # Recursively process all dictionary values for key, value in schema.items(): - _remove_json_schema_refs(value, max_depth - 1) + remove_json_schema_refs(value, max_depth - 1) elif isinstance(schema, list): # Recursively process all items in the list for item in schema: - _remove_json_schema_refs(item, max_depth - 1) + remove_json_schema_refs(item, max_depth - 1) return schema +_remove_json_schema_refs = remove_json_schema_refs + + def _remove_unsupported_params(non_default_params: dict, supported_openai_params: list[str] | None) -> dict: """ Remove unsupported params from non_default_params @@ -4899,7 +4931,7 @@ def get_optional_params( drop_params=bool(drop_params), ) elif custom_llm_provider == "bedrock_mantle": - optional_params = ProviderConfigManager._get_bedrock_mantle_config(model).map_openai_params( + optional_params = ProviderConfigManager.get_bedrock_mantle_config(model).map_openai_params( non_default_params=non_default_params, optional_params=optional_params, model=model, @@ -5019,7 +5051,7 @@ def get_optional_params( ) if _print_verbose_is_active(): print_verbose(f"Final returned optional params: {redact_credentials_in_payload(optional_params)}") - optional_params = _apply_openai_param_overrides( + optional_params = apply_openai_param_overrides( optional_params=optional_params, non_default_params=non_default_params, allowed_openai_params=allowed_openai_params, @@ -5069,8 +5101,8 @@ def add_provider_specific_params_to_optional_params( ) processed_extra_body: Final = {k: v for k, v in initial_extra_body.items() if k not in dropped_keys} - _ensure_extra_body_is_safe: Final = getattr(sys.modules[__name__], "_ensure_extra_body_is_safe") - optional_params["extra_body"] = _ensure_extra_body_is_safe(extra_body=processed_extra_body) + ensure_extra_body_is_safe: Final = getattr(sys.modules[__name__], "ensure_extra_body_is_safe") + optional_params["extra_body"] = ensure_extra_body_is_safe(extra_body=processed_extra_body) else: for k in passed_params: if k not in openai_params and passed_params[k] is not None: @@ -5080,7 +5112,11 @@ def add_provider_specific_params_to_optional_params( return optional_params -def _apply_openai_param_overrides(optional_params: dict, non_default_params: dict, allowed_openai_params: list): +def apply_openai_param_overrides( + optional_params: dict[str, object], + non_default_params: dict[str, object], + allowed_openai_params: list[str], +) -> dict[str, object]: """ If user passes in allowed_openai_params, apply them to optional_params @@ -5103,6 +5139,9 @@ def _apply_openai_param_overrides(optional_params: dict, non_default_params: dic return optional_params +_apply_openai_param_overrides = apply_openai_param_overrides + + PROVIDER_UNVALIDATED_PARAMS: Final = frozenset({"user", "stream_options", "stream", "max_retries"}) @@ -5178,32 +5217,40 @@ def calculate_max_parallel_requests( return None -def _get_deployment_order(deployment: dict | Any) -> int | None: +def get_deployment_order(deployment: Mapping[str, object]) -> int | None: """ Returns the routing order for a deployment. Checks litellm_params first (static config), then model_info (dynamic/team models added via API where order lives in model_info, not litellm_params). """ - order = deployment.get("litellm_params", {}).get("order") - if order is None: - order = deployment.get("model_info", {}).get("order") - return order + litellm_params: Final = cast(Mapping[str, object], deployment.get("litellm_params", {})) + order: Final = litellm_params.get("order") + resolved_order: Final = ( + order if order is not None else cast(Mapping[str, object], deployment.get("model_info", {})).get("order") + ) + return cast(int | None, resolved_order) -def get_order_filtered_deployments(healthy_deployments: list[dict], target_order: int | None = None) -> list: +_get_deployment_order = get_deployment_order + + +def get_order_filtered_deployments( + healthy_deployments: list[dict[str, object]], + target_order: int | None = None, +) -> list[dict[str, object]]: if target_order is not None: - return [d for d in healthy_deployments if _get_deployment_order(d) == target_order] + return [d for d in healthy_deployments if get_deployment_order(d) == target_order] # Default: pick min order group _valid_orders: Final[list[int]] = [ - o for deployment in healthy_deployments for o in [_get_deployment_order(deployment)] if o is not None + o for deployment in healthy_deployments for o in [get_deployment_order(deployment)] if o is not None ] min_order: Final[int | None] = min(_valid_orders) if _valid_orders else None if min_order is not None: filtered_deployments: Final = [ - deployment for deployment in healthy_deployments if _get_deployment_order(deployment) == min_order + deployment for deployment in healthy_deployments if get_deployment_order(deployment) == min_order ] return filtered_deployments @@ -5373,12 +5420,15 @@ def get_first_chars_messages(kwargs: dict) -> str: return "" -def _count_characters(text: str) -> int: +def count_characters(text: str) -> int: # Remove white spaces and count characters filtered_text: Final = "".join(char for char in text if not char.isspace()) return len(filtered_text) +_count_characters = count_characters + + def get_response_string(response_obj: ModelResponse | ModelResponseStream) -> str: # Handle Responses API streaming events if hasattr(response_obj, "type") and hasattr(response_obj, "response"): @@ -5578,7 +5628,7 @@ def _invalidate_model_cost_lowercase_map() -> None: # Clear LRU caches that depend on model_cost data _cached_get_model_info.cache_clear() - _cached_get_model_info_helper.cache_clear() + cached_get_model_info_helper.cache_clear() def _rebuild_model_cost_lowercase_map() -> dict[str, str]: @@ -5630,7 +5680,7 @@ def _handle_new_key_with_scan( return None -def _get_model_cost_key(potential_key: str) -> str | None: +def get_model_cost_key(potential_key: str) -> str | None: """ Get the actual key from model_cost, with case-insensitive fallback. @@ -5671,6 +5721,9 @@ def _get_model_cost_key(potential_key: str) -> str | None: return None +_get_model_cost_key = get_model_cost_key + + def _get_model_info_from_model_cost(key: str) -> dict[str, Any]: return litellm.model_cost[key] @@ -5736,7 +5789,7 @@ class PotentialModelNamesAndCustomLLMProvider(TypedDict): def _first_registered_match( candidates: Sequence[str], custom_llm_provider: str | None ) -> tuple[str | None, dict[str, Any] | None]: - registered_keys: Final = (key for key in map(_get_model_cost_key, candidates) if key is not None) + registered_keys: Final = (key for key in map(get_model_cost_key, candidates) if key is not None) entries: Final = ((key, _get_model_info_from_model_cost(key=key)) for key in registered_keys) matches: Final = ( (key, info) @@ -5773,7 +5826,7 @@ def _get_model_info_from_generalization( potential_model_names["stripped_model_name"], potential_model_names["provider_prefixed_model_name"], ) - if any(_get_model_cost_key(candidate) is not None for candidate in candidates): + if any(get_model_cost_key(candidate) is not None for candidate in candidates): return None for candidate in candidates: generalized_info = match_capability_generalizations(candidate) @@ -5791,7 +5844,7 @@ def _strip_mantle_region_prefix(model: str) -> str: return split_mantle_region_prefix(model)[1] -def _get_potential_model_names(model: str, custom_llm_provider: str | None) -> PotentialModelNamesAndCustomLLMProvider: +def get_potential_model_names(model: str, custom_llm_provider: str | None) -> PotentialModelNamesAndCustomLLMProvider: if custom_llm_provider is None: # Get custom_llm_provider try: @@ -5857,6 +5910,9 @@ def _get_potential_model_names(model: str, custom_llm_provider: str | None) -> P ) +_get_potential_model_names = get_potential_model_names + + def _get_max_position_embeddings(model_name: str) -> int | None: # Construct the URL for the config.json file config_url: Final = f"https://huggingface.co/{model_name}/raw/main/config.json" @@ -5881,7 +5937,7 @@ def _get_max_position_embeddings(model_name: str) -> int | None: @lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE) -def _cached_get_model_info_helper( +def cached_get_model_info_helper( model: str, custom_llm_provider: str | None, api_base: str | None = None, @@ -5891,13 +5947,16 @@ def _cached_get_model_info_helper( Speed Optimization to hit high RPS """ - return _get_model_info_helper( + return get_model_info_helper( model=model, custom_llm_provider=custom_llm_provider, api_base=api_base, ) +_cached_get_model_info_helper = cached_get_model_info_helper + + def get_provider_info(model: str, custom_llm_provider: str | None) -> ProviderSpecificModelInfo | None: ## PROVIDER-SPECIFIC INFORMATION # if custom_llm_provider == "predibase": @@ -5923,7 +5982,7 @@ def _is_potential_model_name_in_model_cost( Check if the potential model name is in the model cost (case-insensitive). """ return any( - _get_model_cost_key(str(potential_model_name)) is not None + get_model_cost_key(str(potential_model_name)) is not None for potential_model_name in potential_model_names.values() ) @@ -5938,7 +5997,7 @@ def _model_not_mapped_message(model: str, custom_llm_provider: str | None) -> st ) -def _get_model_info_helper( +def get_model_info_helper( model: str, custom_llm_provider: str | None = None, api_base: str | None = None, @@ -5963,7 +6022,7 @@ def _get_model_info_helper( ): model = model + "@latest" ########################## - potential_model_names: Final = _get_potential_model_names( + potential_model_names: Final = get_potential_model_names( model=model, custom_llm_provider=custom_llm_provider or declared_authenticating_provider(model) ) @@ -6343,6 +6402,9 @@ def _get_model_info_helper( raise Exception(_model_not_mapped_message(model, custom_llm_provider)) +_get_model_info_helper = get_model_info_helper + + def _build_model_info( model: str, custom_llm_provider: str | None = None, @@ -6351,7 +6413,7 @@ def _build_model_info( ) -> ModelInfo: supported_openai_params = litellm.get_supported_openai_params(model=model, custom_llm_provider=custom_llm_provider) - _model_info: Final = _get_model_info_helper( + _model_info: Final = get_model_info_helper( model=model, custom_llm_provider=custom_llm_provider, api_base=api_base, @@ -7131,7 +7193,7 @@ def check_valid_key(model: str, api_key: str): return False -def _should_retry(status_code: int): +def should_retry(status_code: int) -> bool: """ Retries on 408, 409, 429 and 500 errors. @@ -7160,6 +7222,9 @@ def _should_retry(status_code: int): return False +_should_retry = should_retry + + def _get_retry_after_from_exception_header( response_headers: httpx.Headers | None = None, ): @@ -7194,7 +7259,7 @@ def _get_retry_after_from_exception_header( retry_after = -1 -def _calculate_retry_after( +def calculate_retry_after( remaining_retries: int, max_retries: int, response_headers: httpx.Headers | None = None, @@ -7220,6 +7285,9 @@ def _calculate_retry_after( return sleep_seconds + jitter +_calculate_retry_after = calculate_retry_after + + # custom prompt helper function def register_prompt_template( model: str, @@ -7872,7 +7940,7 @@ def get_valid_models( def print_args_passed_to_litellm(original_function, args, kwargs): - if not _is_debugging_on(): + if not is_debugging_on(): return try: # we've already printed this for acompletion, don't print for completion @@ -7922,32 +7990,33 @@ def get_logging_id(start_time, response_obj): return None -def _get_base_model_from_metadata(model_call_details=None): +def get_base_model_from_metadata(model_call_details: Mapping[str, object] | None = None) -> str | None: if model_call_details is None: return None - litellm_params: Final = model_call_details.get("litellm_params", {}) + litellm_params: Final = cast(Mapping[str, object], model_call_details.get("litellm_params", {})) if litellm_params is not None: - _base_model: Final = litellm_params.get("base_model", None) - if _base_model is not None: - return _base_model - metadata: Final = litellm_params.get("metadata") or {} - - _get_base_model_from_litellm_call_metadata: _BaseModelFromMetadataGetter = getattr( - sys.modules[__name__], "_get_base_model_from_litellm_call_metadata" + base_model: Final = litellm_params.get("base_model", None) + if base_model is not None: + return cast(str | None, base_model) + metadata: Final = cast(Mapping[str, object], litellm_params.get("metadata") or {}) + base_model_metadata_getter: Final[_BaseModelFromMetadataGetter] = getattr( + sys.modules[__name__], "get_base_model_from_litellm_call_metadata" ) - base_model_from_metadata: Final = _get_base_model_from_litellm_call_metadata(metadata=metadata) + base_model_from_metadata: Final = base_model_metadata_getter(metadata=metadata) if base_model_from_metadata is not None: return base_model_from_metadata - # Also check litellm_metadata (used by Responses API and other generic API calls) - litellm_metadata: Final = litellm_params.get("litellm_metadata", {}) - _get_base_model_from_litellm_call_metadata = getattr( - sys.modules[__name__], "_get_base_model_from_litellm_call_metadata" + litellm_metadata: Final = cast(Mapping[str, object], litellm_params.get("litellm_metadata", {})) + litellm_metadata_getter: Final[_BaseModelFromMetadataGetter] = getattr( + sys.modules[__name__], "get_base_model_from_litellm_call_metadata" ) - return _get_base_model_from_litellm_call_metadata(metadata=litellm_metadata) + return litellm_metadata_getter(metadata=litellm_metadata) return None +_get_base_model_from_metadata = get_base_model_from_metadata + + class ModelResponseIterator: def __init__(self, model_response: ModelResponse, convert_to_delta: bool = False): if convert_to_delta is True: @@ -8381,7 +8450,7 @@ def validate_openai_optional_params(stop: str | list[str] | None = None, **kwarg @lru_cache(maxsize=1) -def _get_bundled_model_cost_map() -> dict[str, Any]: +def get_bundled_model_cost_map() -> dict[str, Any]: try: model_cost_path: Final = resources.files("litellm").joinpath("model_prices_and_context_window_backup.json") return json.loads(model_cost_path.read_text()) @@ -8389,6 +8458,9 @@ def _get_bundled_model_cost_map() -> dict[str, Any]: return {} +_get_bundled_model_cost_map = get_bundled_model_cost_map + + def _get_model_cost_entry_for_provider_config( model: str, provider: LlmProviders, @@ -8399,7 +8471,7 @@ def _get_model_cost_entry_for_provider_config( if model_info is not None: return model_info - bundled_model_cost: Final = _get_bundled_model_cost_map() + bundled_model_cost: Final = get_bundled_model_cost_map() for model_key in candidate_keys: model_info = bundled_model_cost.get(model_key) if model_info is not None: @@ -8451,7 +8523,7 @@ class ProviderConfigManager: LlmProviders.DEEPSEEK: (lambda: litellm.DeepSeekChatConfig(), False), LlmProviders.TENCENT: (lambda: litellm.TencentChatConfig(), False), LlmProviders.GROQ: (lambda: litellm.GroqChatConfig(), False), - LlmProviders.BEDROCK_MANTLE: (ProviderConfigManager._get_bedrock_mantle_config, True), + LlmProviders.BEDROCK_MANTLE: (ProviderConfigManager.get_bedrock_mantle_config, True), LlmProviders.A2A: (lambda: litellm.A2AConfig(), False), LlmProviders.BYTEZ: (lambda: litellm.BytezChatConfig(), False), LlmProviders.DATABRICKS: (lambda: litellm.DatabricksConfig(), False), @@ -8640,11 +8712,13 @@ class ProviderConfigManager: return get_bedrock_chat_config(model=model) @staticmethod - def _get_bedrock_mantle_config(model: str) -> BaseConfig: + def get_bedrock_mantle_config(model: str) -> BaseConfig: from litellm.llms.bedrock_mantle.chat.claude_transformation import bedrock_mantle_chat_config return bedrock_mantle_chat_config(model) + _get_bedrock_mantle_config = get_bedrock_mantle_config + @staticmethod def _get_cohere_config(model: str) -> BaseConfig: """Get Cohere config based on route.""" @@ -10109,7 +10183,7 @@ def is_prompt_caching_valid_prompt( model = custom_llm_provider + "/" + model if min_token_count is None: min_token_count = get_prompt_cache_min_tokens(model=model) - return _get_messages_reach_token_count()( + return get_messages_reach_token_count()( model=model, messages=messages, threshold=min_token_count, @@ -10150,7 +10224,7 @@ def extract_duration_from_srt_or_vtt(srt_or_vtt_content: str) -> float | None: return max(durations) if durations else None -def _add_path_to_api_base(api_base: str, ending_path: str) -> str: +def add_path_to_api_base(api_base: str, ending_path: str) -> str: """ Adds an ending path to an API base URL while preventing duplicate path segments. @@ -10188,6 +10262,9 @@ def _add_path_to_api_base(api_base: str, ending_path: str) -> str: return str(modified_url.copy_with(params=original_url.params)) +_add_path_to_api_base = add_path_to_api_base + + def get_standard_openai_params(params: Mapping[str, object]) -> dict: return {k: v for k, v in params.items() if k in litellm.OPENAI_CHAT_COMPLETION_PARAMS and v is not None} @@ -10383,9 +10460,9 @@ def should_run_mock_completion( def __getattr__(name: str) -> Any: """Lazy import handler for utils module with cached registry for improved performance.""" # Use cached registry from _lazy_imports instead of importing tuples every time - from litellm._lazy_imports import _get_lazy_import_registry + from litellm._lazy_imports import get_lazy_import_registry - registry: Final = _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: diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index c0cb3796934..b5c19cddee2 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -86,7 +86,7 @@ class VectorStoreIndexRegistry: ######################################################### @staticmethod - async def _get_vector_store_indexes_from_db( + async def get_vector_store_indexes_from_db( prisma_client: PrismaClient | None, ) -> list[LiteLLM_ManagedVectorStoreIndex]: """ @@ -105,6 +105,8 @@ class VectorStoreIndexRegistry: vector_stores_from_db.append(_litellm_managed_vector_store) return vector_stores_from_db + _get_vector_store_indexes_from_db = get_vector_store_indexes_from_db + class VectorStoreRegistry: def __init__(self, vector_stores: list[LiteLLM_ManagedVectorStore] = []): @@ -242,7 +244,7 @@ class VectorStoreRegistry: # Fall back to database if not found in memory if prisma_client is not None: try: - vector_stores_from_db: Final = await self._get_vector_stores_from_db(prisma_client=prisma_client) + vector_stores_from_db: Final = await self.get_vector_stores_from_db(prisma_client=prisma_client) for db_vector_store in vector_stores_from_db: if db_vector_store.get("vector_store_id") == vector_store_id: # Add to in-memory registry for future use @@ -499,7 +501,7 @@ class VectorStoreRegistry: ######################################################### @staticmethod - async def _get_vector_stores_from_db( + async def get_vector_stores_from_db( prisma_client: PrismaClient | None, ) -> list[LiteLLM_ManagedVectorStore]: """ @@ -516,6 +518,8 @@ class VectorStoreRegistry: vector_stores_from_db.append(_litellm_managed_vector_store) return vector_stores_from_db + _get_vector_stores_from_db = get_vector_stores_from_db + def get_credentials_for_vector_store(self, vector_store_id: str) -> dict[str, object]: """ Get the credentials for a vector store diff --git a/tests/__init__.py b/tests/__init__.py index 9360c76fa85..3ccc7a98a20 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -1 +1 @@ -# This file makes the tests directory a Python package +# This file makes the tests directory a Python package diff --git a/tests/agent_tests/local_only_agent_tests/test_a2a.py b/tests/agent_tests/local_only_agent_tests/test_a2a.py index e2e73808b95..eb0c4ddfc50 100644 --- a/tests/agent_tests/local_only_agent_tests/test_a2a.py +++ b/tests/agent_tests/local_only_agent_tests/test_a2a.py @@ -25,7 +25,7 @@ async def test_asend_message_with_client_decorator(): Test asend_message standalone function with @client decorator. This tests the LiteLLM logging integration. """ - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.a2a_protocol import asend_message, create_a2a_client # Create the A2A client first @@ -191,7 +191,7 @@ async def test_pydantic_ai_non_streaming(): Pydantic AI agents follow A2A protocol but don't support streaming. This test validates non-streaming requests work correctly. """ - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.a2a_protocol import asend_message # Build the request @@ -272,7 +272,7 @@ async def test_pydantic_ai_fake_streaming(): This test validates that fake streaming works by converting non-streaming responses into streaming chunks. """ - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.a2a_protocol import asend_message_streaming # Build the request diff --git a/tests/agent_tests/local_only_agent_tests/test_a2a_completion_bridge.py b/tests/agent_tests/local_only_agent_tests/test_a2a_completion_bridge.py index ff7e9da0368..fb2f19bde2a 100644 --- a/tests/agent_tests/local_only_agent_tests/test_a2a_completion_bridge.py +++ b/tests/agent_tests/local_only_agent_tests/test_a2a_completion_bridge.py @@ -26,7 +26,7 @@ async def test_a2a_completion_bridge_non_streaming(): """ from litellm.a2a_protocol import asend_message - litellm._turn_on_debug() + litellm.turn_on_debug() send_message_payload = { "message": { @@ -77,7 +77,7 @@ async def test_a2a_completion_bridge_streaming(): """ from litellm.a2a_protocol import asend_message_streaming - litellm._turn_on_debug() + litellm.turn_on_debug() send_message_payload = { "message": { @@ -162,7 +162,7 @@ async def test_a2a_completion_bridge_bedrock_agentcore(): """ from litellm.a2a_protocol import asend_message_streaming - litellm._turn_on_debug() + litellm.turn_on_debug() # Bedrock AgentCore ARN (streaming-capable runtime) agentcore_arn = "arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC" @@ -227,7 +227,7 @@ async def test_vertex_agent_engine_non_streaming(): Uses the Reasoning Engine resource ID to call a hosted agent. """ - litellm._turn_on_debug() + litellm.turn_on_debug() # Call via litellm.acompletion with vertex_ai/agent_engine/ prefix response = await litellm.acompletion( @@ -257,7 +257,7 @@ async def test_vertex_agent_engine_streaming(): Uses the Reasoning Engine resource ID to call a hosted agent with streaming. """ - # litellm._turn_on_debug() + # litellm.turn_on_debug() # Call via litellm.acompletion with streaming response = await litellm.acompletion( diff --git a/tests/audio_tests/test_audio_speech.py b/tests/audio_tests/test_audio_speech.py index 77e7fab3f00..e6a9de5a938 100644 --- a/tests/audio_tests/test_audio_speech.py +++ b/tests/audio_tests/test_audio_speech.py @@ -22,7 +22,7 @@ import litellm async def _run_audio_speech_litellm(sync_mode, model, api_base, api_key): - litellm._turn_on_debug() + litellm.turn_on_debug() speech_file_path = Path(__file__).parent / "speech.mp3" if sync_mode: @@ -326,7 +326,7 @@ async def test_azure_ava_tts_async(): """ Test Azure AVA (Cognitive Services) Text-to-Speech with real API request. """ - litellm._turn_on_debug() + litellm.turn_on_debug() api_key = os.getenv("AZURE_TTS_API_KEY") api_base = os.getenv("AZURE_TTS_API_BASE") @@ -380,7 +380,7 @@ async def test_runwayml_tts_async(): """ Test RunwayML Text-to-Speech with real API request. """ - litellm._turn_on_debug() + litellm.turn_on_debug() api_key = os.getenv("RUNWAYML_API_KEY") api_base = os.getenv("RUNWAYML_API_BASE") diff --git a/tests/batches_tests/test_batch_rate_limits.py b/tests/batches_tests/test_batch_rate_limits.py index 7cbcfc1aeb1..daf1bddf042 100644 --- a/tests/batches_tests/test_batch_rate_limits.py +++ b/tests/batches_tests/test_batch_rate_limits.py @@ -73,7 +73,7 @@ async def test_batch_rate_limits(): Integration test for batch rate limits with real OpenAI API calls. Tests the full flow: file creation -> token counting -> cleanup """ - litellm._turn_on_debug() + litellm.turn_on_debug() CUSTOM_LLM_PROVIDER = "openai" BATCH_LIMITER = _build_batch_limiter() @@ -912,9 +912,9 @@ async def test_batch_logging_azure_credentials_regression(): """ from unittest.mock import AsyncMock, MagicMock, patch from litellm.batches.batch_utils import ( - _extract_file_access_credentials, + extract_file_access_credentials, _fetch_batch_output_file_content, - _handle_completed_batch, + handle_completed_batch, ) from litellm.types.llms.openai import Batch, HttpxBinaryResponseContent import httpx @@ -960,7 +960,7 @@ async def test_batch_logging_azure_credentials_regression(): # Test 1: Verify _extract_file_access_credentials works correctly print("\n1. Testing credential extraction...") - extracted_creds = _extract_file_access_credentials(azure_credentials) + extracted_creds = extract_file_access_credentials(azure_credentials) assert "api_key" in extracted_creds, "api_key should be extracted" assert ( extracted_creds["api_key"] == "test-azure-key-regression" @@ -1027,7 +1027,7 @@ async def test_batch_logging_azure_credentials_regression(): with patch( "litellm.files.main.afile_content", side_effect=mock_afile_content_tracker ): - result = await _handle_completed_batch( + result = await handle_completed_batch( batch=mock_batch, custom_llm_provider="azure", litellm_params=azure_credentials, @@ -1064,7 +1064,7 @@ async def test_batch_logging_azure_credentials_regression(): "litellm.files.main.afile_content", side_effect=mock_afile_content_tracker ): try: - result = await _handle_completed_batch( + result = await handle_completed_batch( batch=mock_batch, custom_llm_provider="azure", litellm_params=azure_credentials, diff --git a/tests/batches_tests/test_batches_logging_unit_tests.py b/tests/batches_tests/test_batches_logging_unit_tests.py index 73adb391481..ebe1bb22978 100644 --- a/tests/batches_tests/test_batches_logging_unit_tests.py +++ b/tests/batches_tests/test_batches_logging_unit_tests.py @@ -15,7 +15,7 @@ from litellm import create_batch, create_file from litellm._logging import verbose_logger from litellm.batches.batch_utils import ( _aggregate_batch_cost_usage_models, - _get_file_content_as_dictionary, + get_file_content_as_dictionary, _get_batch_job_usage_from_response_body, _get_response_from_batch_job_output_file, _batch_response_was_successful, @@ -123,7 +123,7 @@ def sample_file_content_dict(): def test_get_file_content_as_dictionary(sample_file_content): - result = _get_file_content_as_dictionary(sample_file_content) + result = get_file_content_as_dictionary(sample_file_content) assert len(result) == 2 assert result[0]["id"] == "batch_req_6769ca596b38819093d7ae9f522de924" assert result[0]["custom_id"] == "request-1" @@ -234,7 +234,7 @@ async def test_batch_retrieve_cost_tracking_with_completed_batch_no_explicit_cos expected_models = ["gpt-5-mini"] with patch( - "litellm.litellm_core_utils.litellm_logging._handle_completed_batch", + "litellm.litellm_core_utils.litellm_logging.handle_completed_batch", new=AsyncMock( return_value=BatchCostUsageResult( cost=expected_cost, @@ -272,7 +272,7 @@ async def test_handle_completed_batch_computes_real_cost_from_output_file( the function the retrieve handler invokes on completion; a dropped output line, a wrong token sum, or mispriced model fails this test. """ - from litellm.batches.batch_utils import _handle_completed_batch + from litellm.batches.batch_utils import handle_completed_batch from litellm.types.utils import LiteLLMBatch batch = LiteLLMBatch( @@ -293,7 +293,7 @@ async def test_handle_completed_batch_computes_real_cost_from_output_file( "litellm.batches.batch_utils._fetch_batch_output_file_content", new=AsyncMock(return_value=sample_file_content_bytes), ): - result = await _handle_completed_batch( + result = await handle_completed_batch( batch=batch, custom_llm_provider="openai" ) @@ -382,7 +382,7 @@ async def test_batch_retrieve_cost_tracking_with_explicit_cost_data(): explicit_models = ["gpt-5-mini", "gpt-5.5"] with patch( - "litellm.litellm_core_utils.litellm_logging._handle_completed_batch", + "litellm.litellm_core_utils.litellm_logging.handle_completed_batch", new=AsyncMock(), ) as mock_handle_batch: # Call async_success_handler with explicit cost data @@ -517,7 +517,7 @@ async def test_batch_retrieve_cost_tracking_with_unified_file_id_incomplete_batc logging_obj.custom_llm_provider = "openai" with patch( - "litellm.litellm_core_utils.litellm_logging._handle_completed_batch", + "litellm.litellm_core_utils.litellm_logging.handle_completed_batch", new=AsyncMock(), ) as mock_handle_batch: # Call async_success_handler with in_progress batch (unified file ID) @@ -606,7 +606,7 @@ async def test_batch_retrieve_cost_tracking_with_partial_explicit_data(): from litellm.batches.batch_utils import BatchCostUsageResult with patch( - "litellm.litellm_core_utils.litellm_logging._handle_completed_batch", + "litellm.litellm_core_utils.litellm_logging.handle_completed_batch", new=AsyncMock( return_value=BatchCostUsageResult( cost=expected_cost, diff --git a/tests/batches_tests/test_bedrock_files_and_batches.py b/tests/batches_tests/test_bedrock_files_and_batches.py index 336fd7dd953..0c66328f3bd 100644 --- a/tests/batches_tests/test_bedrock_files_and_batches.py +++ b/tests/batches_tests/test_bedrock_files_and_batches.py @@ -121,7 +121,7 @@ async def test_async_create_file(): 2. Create Batch Request 3. Retrieve the specific batch """ - litellm._turn_on_debug() + litellm.turn_on_debug() print("Testing async create batch") file_name = "bedrock_batch_completions.jsonl" @@ -162,7 +162,7 @@ async def test_async_file_and_batch(): """ Test file retrieval """ - litellm._turn_on_debug() + litellm.turn_on_debug() file_name = "bedrock_batch_completions.jsonl" _current_dir = os.path.dirname(os.path.abspath(__file__)) file_path = os.path.join(_current_dir, file_name) diff --git a/tests/batches_tests/test_manus_files_all_methods.py b/tests/batches_tests/test_manus_files_all_methods.py index 39311441f59..15ec9378db9 100644 --- a/tests/batches_tests/test_manus_files_all_methods.py +++ b/tests/batches_tests/test_manus_files_all_methods.py @@ -12,7 +12,7 @@ async def test_manus_files_api_e2e_all_methods(): """ E2E test for Manus Files API: create, retrieve, list, delete. """ - litellm._turn_on_debug() + litellm.turn_on_debug() api_key = os.getenv("MANUS_API_KEY") if api_key is None: diff --git a/tests/batches_tests/test_openai_batches_and_files.py b/tests/batches_tests/test_openai_batches_and_files.py index edb64ccb715..758df71dabc 100644 --- a/tests/batches_tests/test_openai_batches_and_files.py +++ b/tests/batches_tests/test_openai_batches_and_files.py @@ -205,7 +205,7 @@ async def test_async_create_batch(provider, tmp_path): 2. Create Batch Request 3. Retrieve the specific batch """ - litellm._turn_on_debug() + litellm.turn_on_debug() print("Testing async create batch") litellm.logging_callback_manager._reset_all_callbacks() @@ -438,7 +438,7 @@ async def test_avertex_batch_prediction(monkeypatch): ) as mock_gcs_upload, ): litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() file_name = "vertex_batch_completions.jsonl" _current_dir = os.path.dirname(os.path.abspath(__file__)) file_path = os.path.join(_current_dir, file_name) @@ -591,7 +591,7 @@ async def test_delete_batch_output_file(): - The output file can be deleted without validation errors - The file_object is fetched and stored with proper metadata instead of None """ - litellm._turn_on_debug() + litellm.turn_on_debug() print("Testing delete batch output file") file_name = "openai_batch_completions.jsonl" diff --git a/tests/benchmarks/test_benchmarks.py b/tests/benchmarks/test_benchmarks.py index 59b3e0b6d5c..c7d815ce331 100644 --- a/tests/benchmarks/test_benchmarks.py +++ b/tests/benchmarks/test_benchmarks.py @@ -201,13 +201,13 @@ def test_cost_per_token_anthropic(): @pytest.mark.benchmark def test_get_model_cost_key_exact_match(): """Benchmark model cost key lookup with an exact match.""" - litellm.utils._get_model_cost_key("gpt-4o") + litellm.utils.get_model_cost_key("gpt-4o") @pytest.mark.benchmark def test_get_model_cost_key_case_insensitive(): """Benchmark model cost key lookup with case-insensitive fallback.""" - litellm.utils._get_model_cost_key("GPT-4o") + litellm.utils.get_model_cost_key("GPT-4o") # --------------------------------------------------------------------------- diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 863e8befcf9..ca0cb7bd132 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -3,8 +3,8 @@ import os IGNORE_FUNCTIONS = [ "_format_type", - "_remove_additional_properties", - "_remove_strict_from_schema", + "remove_additional_properties", + "remove_strict_from_schema", "filter_schema_fields", "text_completion", "_check_for_os_environ_vars", @@ -25,7 +25,7 @@ IGNORE_FUNCTIONS = [ "filter_value_from_dict", # max depth set. "normalize_json_schema_types", # max depth set. "_extract_fields_recursive", # max depth set. - "_remove_json_schema_refs", # max depth set., + "remove_json_schema_refs", # max depth set., "_convert_schema_types", # max depth set., "_fix_enum_empty_strings", # max depth set., "get_access_token", # max depth set., @@ -49,7 +49,7 @@ IGNORE_FUNCTIONS = [ "_convert_to_json_serializable_dict", # max depth set (default 20) and circular reference protection to prevent infinite recursion. "dict", # max depth set. _LiteLLMParamsDictView.dict() calls builtin dict(), not itself. "_read_image_bytes", # max depth set. - "_get_masked_values", # max depth set (default 20) to prevent infinite recursion while masking nested sensitive config dicts. + "get_masked_values", # max depth set (default 20) to prevent infinite recursion while masking nested sensitive config dicts. "_redact_sensitive_litellm_params", # max depth set (default 10). "_redact_secret_values_in_obj", # max depth set (default 10, _REDACT_SECRET_MAX_DEPTH); fails closed by returning "REDACTED" at the cap. "_resolve", # OCI: $ref resolver bounded by `resolving_stack` cycle guard. diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index baa46aa5331..533afc6c94f 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -80,6 +80,7 @@ ignored_function_names = [ "chunks", # Property on FallbackAwareAnthropicMessagesStream, so its reads in tests are never an ast.Call "messages", # Property on FallbackAwareAnthropicMessagesStream, so its reads in tests are never an ast.Call "model", # Property on FallbackAwareAnthropicMessagesStream, so its reads in tests are never an ast.Call + "_routing_groups", # Property getter and setter reads are never ast.Call nodes "_request_header", # Tested through Claude Code session routing in test_router.py "_claude_code_session_router_cache_key", # Tested through Claude Code session routing in test_router.py "_delete_claude_code_session_router_binding", # Tested through Redis cleanup failure in test_router.py diff --git a/tests/guardrails_tests/test_javelin_guardrails.py b/tests/guardrails_tests/test_javelin_guardrails.py index a2e7747d657..8ec26b467fe 100644 --- a/tests/guardrails_tests/test_javelin_guardrails.py +++ b/tests/guardrails_tests/test_javelin_guardrails.py @@ -13,7 +13,7 @@ async def test_javelin_guardrail_reject_prompt(): """ Test that the Javelin guardrail raises HTTPException when violations are detected, preventing the request from going to the LLM. """ - # litellm._turn_on_debug() + # litellm.turn_on_debug() guardrail = JavelinGuardrail( guardrail_name="promptinjectiondetection", api_base="https://api-dev.javelin.live", diff --git a/tests/guardrails_tests/test_lakera_v2.py b/tests/guardrails_tests/test_lakera_v2.py index a71759862b2..29b001b4d7d 100644 --- a/tests/guardrails_tests/test_lakera_v2.py +++ b/tests/guardrails_tests/test_lakera_v2.py @@ -18,7 +18,7 @@ from litellm.types.utils import CallTypes as LitellmCallTypes, ModelResponse async def test_lakera_pre_call_hook_for_pii_masking(): """Test for Lakera guardrail pre-call hook for PII masking""" # Setup the guardrail with specific entities config - litellm._turn_on_debug() + litellm.turn_on_debug() lakera_guardrail = LakeraAIGuardrail( api_key="test_key", ) diff --git a/tests/guardrails_tests/test_presidio_pii.py b/tests/guardrails_tests/test_presidio_pii.py index 2960adfcbd3..60c614e88cc 100644 --- a/tests/guardrails_tests/test_presidio_pii.py +++ b/tests/guardrails_tests/test_presidio_pii.py @@ -18,7 +18,7 @@ from litellm.exceptions import BlockedPiiEntityError async def test_presidio_with_blocked_entities(): """Test for Presidio guardrail with blocked entities - requires actual Presidio API""" # Setup the guardrail with specific entities config - BLOCK for credit card - litellm._turn_on_debug() + litellm.turn_on_debug() pii_entities_config = { PiiEntityType.CREDIT_CARD: PiiAction.BLOCK, # This entity should cause a block PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, # This entity should be masked @@ -377,7 +377,7 @@ async def test_presidio_pii_masking_logging_output_only_logged_response_guardrai @pytest.mark.asyncio async def test_presidio_language_configuration(): """Test that presidio_language parameter is properly set and used in analyze requests""" - litellm._turn_on_debug() + litellm.turn_on_debug() # Test with German language using mock testing to avoid API calls presidio_guardrail_de = _OPTIONAL_PresidioPIIMasking( @@ -433,7 +433,7 @@ async def test_presidio_language_configuration(): @pytest.mark.asyncio async def test_presidio_language_configuration_with_per_request_override(): """Test that per-request language configuration overrides the default configured language""" - litellm._turn_on_debug() + litellm.turn_on_debug() # Set up guardrail with German as default language presidio_guardrail = _OPTIONAL_PresidioPIIMasking( diff --git a/tests/guardrails_tests/test_tracing_guardrails.py b/tests/guardrails_tests/test_tracing_guardrails.py index ac85803ba39..11c18969be3 100644 --- a/tests/guardrails_tests/test_tracing_guardrails.py +++ b/tests/guardrails_tests/test_tracing_guardrails.py @@ -209,7 +209,7 @@ async def test_langfuse_trace_includes_guardrail_information(): mock_post.return_value = mock_response with patch("httpx.Client.post", mock_post): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.callbacks = [callback] presidio_guard = _OPTIONAL_PresidioPIIMasking( guardrail_name="presidio_guard", @@ -311,7 +311,7 @@ async def test_bedrock_guardrail_status_blocked(): from litellm.proxy._types import UserAPIKeyAuth from unittest.mock import AsyncMock, MagicMock, patch - litellm._turn_on_debug() + litellm.turn_on_debug() # Setup custom logger to capture standard logging payload test_custom_logger = CustomLoggerForTesting() diff --git a/tests/image_gen_tests/base_image_generation_test.py b/tests/image_gen_tests/base_image_generation_test.py index c50b09d329c..24ef3c6f1b7 100644 --- a/tests/image_gen_tests/base_image_generation_test.py +++ b/tests/image_gen_tests/base_image_generation_test.py @@ -42,7 +42,7 @@ class BaseImageGenTest(ABC): async def test_basic_image_generation(self): """Test basic image generation""" try: - litellm._turn_on_debug() + litellm.turn_on_debug() custom_logger = TestCustomLogger() litellm.logging_callback_manager._reset_all_callbacks() litellm.callbacks = [custom_logger] diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py index ff0cf7e3075..94ced8040a9 100644 --- a/tests/image_gen_tests/test_image_edits.py +++ b/tests/image_gen_tests/test_image_edits.py @@ -64,7 +64,7 @@ class BaseLLMImageEditTest(ABC): """ Test image edit functionality with both sync and async modes. """ - litellm._turn_on_debug() + litellm.turn_on_debug() try: prompt = """ Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. @@ -158,7 +158,7 @@ class TestAzureAIFlux2ImageEdit(BaseLLMImageEditTest): @pytest.mark.flaky(retries=3, delay=2) @pytest.mark.asyncio async def test_openai_image_edit_litellm_router(): - litellm._turn_on_debug() + litellm.turn_on_debug() try: prompt = """ Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. @@ -201,7 +201,7 @@ async def test_openai_image_edit_with_bytesio(): """Test image editing using BytesIO objects instead of file readers""" from litellm import image_edit, aimage_edit - litellm._turn_on_debug() + litellm.turn_on_debug() try: prompt = """ Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. @@ -264,7 +264,7 @@ async def test_azure_image_edit_litellm_sdk(): # Configure the mock to return our response mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() prompt = """ Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. @@ -385,7 +385,7 @@ async def test_openai_image_edit_cost_tracking(): # Configure the mock to return our response mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() prompt = """ Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. @@ -476,7 +476,7 @@ async def test_azure_image_edit_cost_tracking(): # Configure the mock to return our response mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() prompt = """ Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. @@ -535,7 +535,7 @@ async def test_recraft_image_edit_api(): from litellm import aimage_edit import requests - litellm._turn_on_debug() + litellm.turn_on_debug() try: prompt = """ Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. @@ -630,7 +630,7 @@ async def test_multiple_image_edit_with_different_formats(): """Test multiple images editing with different file formats and types""" from litellm import aimage_edit - litellm._turn_on_debug() + litellm.turn_on_debug() try: prompt = "Create a cohesive artistic style across all images" diff --git a/tests/image_gen_tests/test_image_generation.py b/tests/image_gen_tests/test_image_generation.py index 02cee2e8a00..57fb985a747 100644 --- a/tests/image_gen_tests/test_image_generation.py +++ b/tests/image_gen_tests/test_image_generation.py @@ -188,7 +188,7 @@ class TestAimlImageGeneration(BaseImageGenTest): mock_sync_post.return_value = mock_response try: - litellm._turn_on_debug() + litellm.turn_on_debug() custom_logger = TestCustomLogger() litellm.logging_callback_manager._reset_all_callbacks() litellm.callbacks = [custom_logger] diff --git a/tests/integration/spend/test_spend_log_tool_payload_content.py b/tests/integration/spend/test_spend_log_tool_payload_content.py index 556e0b9dbc2..22ed6c74756 100644 --- a/tests/integration/spend/test_spend_log_tool_payload_content.py +++ b/tests/integration/spend/test_spend_log_tool_payload_content.py @@ -91,7 +91,7 @@ def _sse_events(body: str) -> tuple[dict[str, JsonValue], ...]: def _spend_request_id(response_id: str, *, responses_api: bool = False) -> str: if not responses_api: return response_id - decoded: Final = ResponsesAPIRequestUtils._decode_responses_api_response_id(response_id) + decoded: Final = ResponsesAPIRequestUtils.decode_responses_api_response_id(response_id) request_id: Final = decoded.get("response_id") return string_value(request_id) if isinstance(request_id, str) else response_id diff --git a/tests/litellm_utils_tests/test_health_check.py b/tests/litellm_utils_tests/test_health_check.py index 499bcdd5910..bce7bfdf615 100644 --- a/tests/litellm_utils_tests/test_health_check.py +++ b/tests/litellm_utils_tests/test_health_check.py @@ -88,7 +88,7 @@ async def test_azure_img_gen_health_check(): Test Azure image generation health check with retry logic for transient errors. Azure sometimes returns internal server errors which are transient and not something we can control. """ - litellm._turn_on_debug() + litellm.turn_on_debug() max_retries = 3 retry_delay = 1 # Start with 1 second delay @@ -724,7 +724,7 @@ async def test_timeout_does_not_cancel_other_health_checks(): @pytest.mark.asyncio async def test_ahealth_check_ocr(): - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.ahealth_check( model_params={ "model": "mistral/mistral-ocr-latest", diff --git a/tests/litellm_utils_tests/test_logging_callback_manager.py b/tests/litellm_utils_tests/test_logging_callback_manager.py index ebd5b473ebb..4689f58696a 100644 --- a/tests/litellm_utils_tests/test_logging_callback_manager.py +++ b/tests/litellm_utils_tests/test_logging_callback_manager.py @@ -305,7 +305,7 @@ async def test_generic_api_compatible_callbacks_json(): with patch.dict(os.environ, {"SUMOLOGIC_WEBHOOK_URL": test_sumologic_url}): # Test that sumologic callback is recognized from JSON file - result = LoggingCallbackManager._add_custom_callback_generic_api_str( + result = LoggingCallbackManager.add_custom_callback_generic_api_str( "sumologic" ) @@ -346,7 +346,7 @@ async def test_generic_api_compatible_callbacks_json_rubrik(): {"RUBRIK_WEBHOOK_URL": test_rubrik_url, "RUBRIK_API_KEY": test_rubrik_api_key}, ): # Test that rubrik callback is recognized from JSON file - result = LoggingCallbackManager._add_custom_callback_generic_api_str("rubrik") + result = LoggingCallbackManager.add_custom_callback_generic_api_str("rubrik") # Verify a GenericAPILogger instance is returned assert isinstance( @@ -378,7 +378,7 @@ def test_generic_api_compatible_callbacks_json_unknown_callback(): Test that unknown callbacks (not in JSON or callback_settings) are returned unchanged """ # Test with a callback that doesn't exist in the JSON file - result = LoggingCallbackManager._add_custom_callback_generic_api_str( + result = LoggingCallbackManager.add_custom_callback_generic_api_str( "unknown_callback" ) @@ -409,7 +409,7 @@ async def test_generic_api_callback_settings_retry_config(): } try: - result = LoggingCallbackManager._add_custom_callback_generic_api_str( + result = LoggingCallbackManager.add_custom_callback_generic_api_str( callback_name ) diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index 401434c9d36..a0589136996 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -43,7 +43,7 @@ def reset_mock_cache(): # Test 1: Check trimming of normal message def test_basic_trimming(): - litellm._turn_on_debug() + litellm.turn_on_debug() messages = [ { "role": "user", @@ -873,7 +873,7 @@ def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id): ) time.sleep(3) - assert litellm_logging_obj._get_trace_id(service_name="langfuse") is not None + assert litellm_logging_obj.get_trace_id(service_name="langfuse") is not None # langfuse addresses a trace by a 32-hex id, so the id litellm reports back is the # resolved form of whichever source won; that is what the alerting deep link needs @@ -884,7 +884,7 @@ def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id): else: expected_source = litellm_logging_obj.litellm_trace_id - assert litellm_logging_obj._get_trace_id(service_name="langfuse") == resolve_trace_id( + assert litellm_logging_obj.get_trace_id(service_name="langfuse") == resolve_trace_id( expected_source ) @@ -1350,7 +1350,7 @@ def test_is_prompt_caching_enabled_return_default_image_dimensions(): """ mock_token_counter = MagicMock(return_value=False) with patch( - "litellm.utils._get_messages_reach_token_count", + "litellm.utils.get_messages_reach_token_count", return_value=mock_token_counter, ): litellm.utils.is_prompt_caching_valid_prompt( @@ -1434,7 +1434,7 @@ def test_get_valid_models_openai_proxy(monkeypatch): from litellm.utils import get_valid_models import litellm - litellm._turn_on_debug() + litellm.turn_on_debug() monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-9876") monkeypatch.setenv("LITELLM_PROXY_API_BASE", "https://litellm-api.up.railway.app/") @@ -1469,7 +1469,7 @@ def test_get_valid_models_fireworks_ai(monkeypatch): from litellm.utils import get_valid_models import litellm - litellm._turn_on_debug() + litellm.turn_on_debug() monkeypatch.setenv("FIREWORKS_API_KEY", "sk-9876") monkeypatch.setenv("FIREWORKS_ACCOUNT_ID", "1234") @@ -1587,12 +1587,12 @@ def test_get_num_retries(num_retries): def test_add_custom_logger_callback_to_specific_event(monkeypatch): - from litellm.utils import _add_custom_logger_callback_to_specific_event + from litellm.utils import add_custom_logger_callback_to_specific_event monkeypatch.setattr(litellm, "success_callback", []) monkeypatch.setattr(litellm, "failure_callback", []) - _add_custom_logger_callback_to_specific_event("langfuse", "success") + add_custom_logger_callback_to_specific_event("langfuse", "success") assert len(litellm.success_callback) == 1 assert len(litellm.failure_callback) == 0 @@ -2270,13 +2270,13 @@ def test_get_base_model_from_metadata(): Related issue: https://github.com/BerriAI/litellm/issues/16772 """ - from litellm.utils import _get_base_model_from_metadata + from litellm.utils import get_base_model_from_metadata # Test 1: base_model in metadata (Chat Completions API pattern) model_call_details_with_metadata = { "litellm_params": {"metadata": {"model_info": {"base_model": "azure/gpt-5.5"}}} } - result = _get_base_model_from_metadata(model_call_details_with_metadata) + result = get_base_model_from_metadata(model_call_details_with_metadata) assert result == "azure/gpt-5.5", f"Expected 'azure/gpt-5.5', got {result}" # Test 2: base_model in litellm_metadata (Responses API and generic API calls pattern) @@ -2285,14 +2285,14 @@ def test_get_base_model_from_metadata(): "litellm_metadata": {"model_info": {"base_model": "azure/gpt-5-mini"}} } } - result = _get_base_model_from_metadata(model_call_details_with_litellm_metadata) + result = get_base_model_from_metadata(model_call_details_with_litellm_metadata) assert result == "azure/gpt-5-mini", f"Expected 'azure/gpt-5-mini', got {result}" # Test 3: base_model in litellm_params (direct base_model) model_call_details_with_direct_base_model = { "litellm_params": {"base_model": "azure/gpt-5-mini"} } - result = _get_base_model_from_metadata(model_call_details_with_direct_base_model) + result = get_base_model_from_metadata(model_call_details_with_direct_base_model) assert ( result == "azure/gpt-5-mini" ), f"Expected 'azure/gpt-5-mini', got {result}" @@ -2306,16 +2306,16 @@ def test_get_base_model_from_metadata(): }, } } - result = _get_base_model_from_metadata(model_call_details_with_both) + result = get_base_model_from_metadata(model_call_details_with_both) assert ( result == "azure/gpt-4-from-metadata" ), f"Expected metadata to take precedence, got {result}" # Test 5: No base_model present model_call_details_without_base_model = {"litellm_params": {"metadata": {}}} - result = _get_base_model_from_metadata(model_call_details_without_base_model) + result = get_base_model_from_metadata(model_call_details_without_base_model) assert result is None, f"Expected None when no base_model present, got {result}" # Test 6: None input - result = _get_base_model_from_metadata(None) + result = get_base_model_from_metadata(None) assert result is None, f"Expected None for None input, got {result}" diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py index f6309ce6990..8f3867a9201 100644 --- a/tests/llm_responses_api_testing/base_responses_api.py +++ b/tests/llm_responses_api_testing/base_responses_api.py @@ -112,7 +112,7 @@ class BaseResponsesAPITest(ABC): @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_basic_openai_responses_api(self, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True base_completion_call_args = self.get_base_completion_call_args() try: @@ -139,7 +139,7 @@ class BaseResponsesAPITest(ABC): @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=2) async def test_basic_openai_responses_api_streaming(self, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() # Enable cost calculation for streaming usage litellm.include_cost_in_streaming_usage = True base_completion_call_args = self.get_base_completion_call_args() @@ -232,7 +232,7 @@ class BaseResponsesAPITest(ABC): @pytest.mark.parametrize("sync_mode", [False, True]) @pytest.mark.asyncio async def test_basic_openai_responses_delete_endpoint(self, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True base_completion_call_args = self.get_base_completion_call_args() if sync_mode: @@ -264,7 +264,7 @@ class BaseResponsesAPITest(ABC): @pytest.mark.flaky(retries=3, delay=2) @pytest.mark.asyncio async def test_basic_openai_responses_streaming_delete_endpoint(self, sync_mode): - # litellm._turn_on_debug() + # litellm.turn_on_debug() # litellm.set_verbose = True base_completion_call_args = self.get_base_completion_call_args() response_id = None @@ -314,7 +314,7 @@ class BaseResponsesAPITest(ABC): @pytest.mark.flaky(retries=3, delay=2) @pytest.mark.asyncio async def test_basic_openai_responses_get_endpoint(self, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True base_completion_call_args = self.get_base_completion_call_args() if sync_mode: @@ -349,7 +349,7 @@ class BaseResponsesAPITest(ABC): @pytest.mark.asyncio async def test_multiturn_responses_api(self): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True try: base_completion_call_args = self.get_base_completion_call_args() @@ -375,7 +375,7 @@ class BaseResponsesAPITest(ABC): @pytest.mark.asyncio async def test_responses_api_with_tool_calls(self): """Test that calls the Responses API with tool calls including function call and output""" - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True base_completion_call_args = self.get_base_completion_call_args() @@ -518,7 +518,7 @@ class BaseResponsesAPITest(ABC): @pytest.mark.asyncio async def test_basic_openai_responses_cancel_endpoint(self, sync_mode): try: - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True base_completion_call_args = self.get_base_completion_call_args() if sync_mode: diff --git a/tests/llm_responses_api_testing/test_anthropic_responses_api.py b/tests/llm_responses_api_testing/test_anthropic_responses_api.py index 8ed85aaa209..8d62e4819e4 100644 --- a/tests/llm_responses_api_testing/test_anthropic_responses_api.py +++ b/tests/llm_responses_api_testing/test_anthropic_responses_api.py @@ -28,7 +28,7 @@ from openai.types.responses.function_tool import FunctionTool class TestAnthropicResponsesAPITest(BaseResponsesAPITest): def get_base_completion_call_args(self): - # litellm._turn_on_debug() + # litellm.turn_on_debug() return { "model": "anthropic/claude-sonnet-4-5", } @@ -53,7 +53,7 @@ class TestAnthropicResponsesAPITest(BaseResponsesAPITest): def test_multiturn_tool_calls(): # Test streaming response with tools for Anthropic - litellm._turn_on_debug() + litellm.turn_on_debug() shell_tool = dict( FunctionTool( type="function", diff --git a/tests/llm_responses_api_testing/test_azure_responses_api.py b/tests/llm_responses_api_testing/test_azure_responses_api.py index 1ec7bafd1ad..2372074866c 100644 --- a/tests/llm_responses_api_testing/test_azure_responses_api.py +++ b/tests/llm_responses_api_testing/test_azure_responses_api.py @@ -36,7 +36,7 @@ async def test_azure_responses_api_preview_api_version(): """ Ensure new azure preview api version is working """ - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.aresponses( model="azure/gpt-5-mini", truncation="auto", diff --git a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py index da37803b64a..7cdee04760c 100644 --- a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py +++ b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py @@ -148,7 +148,7 @@ class TestBaseResponsesAPIStreamingIterator: with patch.object( ResponsesAPIRequestUtils, - "_update_responses_api_response_id_with_model_id", + "update_responses_api_response_id_with_model_id", return_value=updated_response, ) as mock_update_id: # Process the chunk @@ -218,7 +218,7 @@ class TestBaseResponsesAPIStreamingIterator: } with patch.object( - ResponsesAPIRequestUtils, "_update_responses_api_response_id_with_model_id" + ResponsesAPIRequestUtils, "update_responses_api_response_id_with_model_id" ) as mock_update_id: # Process the chunk result = iterator._process_chunk(json.dumps(test_chunk_data)) @@ -584,7 +584,7 @@ class TestBaseResponsesAPIStreamingIterator: with ( patch.object( ResponsesAPIRequestUtils, - "_update_responses_api_response_id_with_model_id", + "update_responses_api_response_id_with_model_id", return_value=mock_responses_api_response, ), patch( @@ -661,7 +661,7 @@ class TestBaseResponsesAPIStreamingIterator: with ( patch.object( ResponsesAPIRequestUtils, - "_update_responses_api_response_id_with_model_id", + "update_responses_api_response_id_with_model_id", return_value=mock_responses_api_response, ), patch("asyncio.create_task") as mock_create_task, diff --git a/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py b/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py index d84e9cc66e3..6cc2a60560b 100644 --- a/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py +++ b/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py @@ -9,7 +9,7 @@ from base_responses_api import BaseResponsesAPITest @pytest.mark.asyncio async def test_basic_google_ai_studio_responses_api_with_tools(): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True request_model = "gemini/gemini-2.5-flash" response = await litellm.aresponses( @@ -28,7 +28,7 @@ async def test_mock_basic_google_ai_studio_responses_api_with_tools(): litellm.acompletion(messages=[{'role': 'user', 'content': 'what is the latest version of supabase python package and when was it released?'}], model='gemini-2.5-flash', tools=[], web_search_options={'search_context_size': 'low', 'user_location': None}) """ # Mock the acompletion function - litellm._turn_on_debug() + litellm.turn_on_debug() mock_response = litellm.ModelResponse( id="test-id", created=1234567890, @@ -276,7 +276,7 @@ async def test_gemini_3_responses_api_streaming_with_thought_signatures(): class TestGoogleAIStudioResponsesAPITest(BaseResponsesAPITest): def get_base_completion_call_args(self): - # litellm._turn_on_debug() + # litellm.turn_on_debug() return {"model": "gemini/gemini-2.5-flash-lite"} async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False): diff --git a/tests/llm_responses_api_testing/test_openai_responses_api.py b/tests/llm_responses_api_testing/test_openai_responses_api.py index 051eb7494b2..60349bb00fc 100644 --- a/tests/llm_responses_api_testing/test_openai_responses_api.py +++ b/tests/llm_responses_api_testing/test_openai_responses_api.py @@ -94,7 +94,7 @@ def validate_standard_logging_payload( @pytest.mark.asyncio def test_basic_openai_responses_api_streaming_with_logging(): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True test_custom_logger = TestCustomLogger() litellm.callbacks = [test_custom_logger] @@ -166,7 +166,7 @@ def validate_responses_match(slp_response, litellm_response): @pytest.mark.asyncio async def test_basic_openai_responses_api_non_streaming_with_logging(): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True test_custom_logger = TestCustomLogger() litellm.callbacks = [test_custom_logger] @@ -204,7 +204,7 @@ async def test_openai_responses_api_returns_headers(sync_mode): Related issue: LiteLLM responses API should return OpenAI headers like chat completions does """ - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True if sync_mode: @@ -459,7 +459,7 @@ def validate_stream_event(event): @pytest.mark.asyncio async def test_openai_responses_api_streaming_validation(sync_mode): """Test that validates each streaming event from the responses API""" - litellm._turn_on_debug() + litellm.turn_on_debug() event_types_seen = set() @@ -499,7 +499,7 @@ async def test_openai_responses_litellm_router(sync_mode): """ Test the OpenAI responses API with LiteLLM Router in both sync and async modes """ - litellm._turn_on_debug() + litellm.turn_on_debug() router = litellm.Router( model_list=[ { @@ -544,7 +544,7 @@ async def test_openai_responses_litellm_router_streaming(sync_mode): """ Test the OpenAI responses API with streaming through LiteLLM Router """ - litellm._turn_on_debug() + litellm.turn_on_debug() router = litellm.Router( model_list=[ { @@ -652,7 +652,7 @@ async def test_openai_responses_litellm_router_no_metadata(): # Configure the mock to return our response mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() router = litellm.Router( model_list=[ { @@ -750,7 +750,7 @@ async def test_openai_responses_litellm_router_with_metadata(): # Configure the mock to return our response mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() router = litellm.Router( model_list=[ { @@ -832,7 +832,7 @@ async def test_openai_responses_litellm_router_with_prompt(): ) as mock_post: mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() router = litellm.Router( model_list=[ { @@ -950,7 +950,7 @@ async def test_openai_o1_pro_response_api(sync_mode): # Configure the mock to return our response mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True # Call o1-pro with max_output_tokens=20 @@ -1047,7 +1047,7 @@ async def test_openai_o1_pro_response_api_streaming(sync_mode): # Configure the mock to return our response mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True # Verify the request was made correctly @@ -1165,7 +1165,7 @@ def test_basic_computer_use_preview_tool_call(): "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", return_value=MockResponse(mock_response, 200), ) as mock_post: - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True # Call the responses API with computer_use_preview tool @@ -1206,7 +1206,7 @@ def test_basic_computer_use_preview_tool_call(): def test_mcp_tools_with_responses_api(): - litellm._turn_on_debug() + litellm.turn_on_debug() MCP_TOOLS = [ { "type": "mcp", @@ -1269,7 +1269,7 @@ def test_mcp_tools_with_responses_api(): @pytest.mark.asyncio async def test_openai_responses_api_field_types(): """Test that specific fields in the response have the correct types""" - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True # Test with store=True @@ -1473,7 +1473,7 @@ async def test_aresponses_service_tier_and_safety_identifier(): # Configure the mock to return our response mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True # Call aresponses with service_tier and safety_identifier @@ -1570,7 +1570,7 @@ async def test_openai_gpt5_reasoning_effort_parameter(): # Configure the mock to return our response mock_post.return_value = MockResponse(mock_response, 200) - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True # Call aresponses with reasoning_effort parameter @@ -1611,7 +1611,7 @@ async def test_openai_responses_api_token_limit_error(): carrying the provider's message. invalid_request_error is a non-retriable client error, so there is no MidStreamFallbackError wrapping. """ - litellm._turn_on_debug() + litellm.turn_on_debug() # Generate text with >400k tokens to trigger token limit error oversized_text = "This is a test sentence. " * 50000 # ~400k tokens @@ -1633,7 +1633,7 @@ async def test_openai_responses_api_token_limit_error(): async def test_openai_streaming_logging(): """Test that OpenAI Responses API streaming logging is working correctly.""" - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.integrations.custom_logger import CustomLogger from litellm.types.utils import Usage @@ -1819,7 +1819,7 @@ async def test_openai_compact_responses_api(sync_mode): This test verifies that the compact_responses endpoint works correctly for compressing conversation history. """ - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True input_messages = [ diff --git a/tests/llm_responses_api_testing/test_responses_hooks.py b/tests/llm_responses_api_testing/test_responses_hooks.py index a86752c0172..28f1b50186f 100644 --- a/tests/llm_responses_api_testing/test_responses_hooks.py +++ b/tests/llm_responses_api_testing/test_responses_hooks.py @@ -57,7 +57,7 @@ class _FakeLoggingObj: async def async_failure_handler(self, *args, **kwargs): self.async_failure_calls += 1 - def _update_completion_start_time(self, completion_start_time): + def update_completion_start_time(self, completion_start_time): self.completion_start_time = completion_start_time self.model_call_details["completion_start_time"] = completion_start_time @@ -377,7 +377,7 @@ def test_process_chunk_completed_response_updates_id_and_usage_cost(monkeypatch) ) logging_obj = _FakeLoggingObj() - logging_obj._response_cost_calculator = MagicMock(return_value=1.23) + logging_obj.response_cost_calculator = MagicMock(return_value=1.23) iterator = ResponsesAPIStreamingIterator( response=httpx.Response(200), model="test-model", @@ -521,7 +521,7 @@ def test_process_chunk_cost_annotation_failure_is_nonfatal(monkeypatch): ) logging_obj = _FakeLoggingObj() - logging_obj._response_cost_calculator = MagicMock(side_effect=RuntimeError("boom")) + logging_obj.response_cost_calculator = MagicMock(side_effect=RuntimeError("boom")) iterator = ResponsesAPIStreamingIterator( response=httpx.Response(200), model="test-model", @@ -584,7 +584,7 @@ async def test_responses_streaming_completed_event_persists_async_cache(): async_set_cache=AsyncMock(), _should_store_result_in_cache=lambda original_function, kwargs: True, ) - logging_obj._llm_caching_handler = caching_handler + logging_obj.llm_caching_handler = caching_handler iterator = ResponsesAPIStreamingIterator( response=httpx.Response(200), @@ -636,7 +636,7 @@ def test_responses_streaming_completed_event_persists_sync_cache(): sync_set_cache=MagicMock(), _should_store_result_in_cache=lambda original_function, kwargs: True, ) - logging_obj._llm_caching_handler = caching_handler + logging_obj.llm_caching_handler = caching_handler iterator = SyncResponsesAPIStreamingIterator( response=httpx.Response(200), @@ -768,9 +768,9 @@ def test_persist_completed_response_to_cache_guard_branches(monkeypatch, scenari response=completed_event.response, ) elif scenario == "missing_caching_handler": - logging_obj._llm_caching_handler = None + logging_obj.llm_caching_handler = None else: - logging_obj._llm_caching_handler = SimpleNamespace( + logging_obj.llm_caching_handler = SimpleNamespace( request_kwargs={ "model": "test-model", "input": "hello", @@ -805,7 +805,7 @@ def test_build_synthetic_response_events_covers_annotations_function_calls_and_r original_include_cost = litellm.include_cost_in_streaming_usage litellm.include_cost_in_streaming_usage = True logging_obj = _FakeLoggingObj() - logging_obj._response_cost_calculator = MagicMock(side_effect=RuntimeError("boom")) + logging_obj.response_cost_calculator = MagicMock(side_effect=RuntimeError("boom")) transformed = ResponsesAPIResponse( id="resp_events", created_at=int(datetime.now().timestamp()), @@ -985,7 +985,7 @@ async def test_cached_responses_stream_async_hit_triggers_success_callbacks( async_add_cache=AsyncMock(), add_cache=MagicMock(), ) - logging_obj._llm_caching_handler = SimpleNamespace( + logging_obj.llm_caching_handler = SimpleNamespace( request_kwargs={"model": "test-model", "input": "hello", "stream": True}, preset_cache_key="responses-stream-cache-key", original_function=litellm.aresponses, @@ -1040,7 +1040,7 @@ def test_cached_responses_stream_sync_hit_triggers_success_callbacks(monkeypatch async_add_cache=AsyncMock(), add_cache=MagicMock(), ) - logging_obj._llm_caching_handler = SimpleNamespace( + logging_obj.llm_caching_handler = SimpleNamespace( request_kwargs={"model": "test-model", "input": "hello", "stream": True}, preset_cache_key="responses-stream-cache-key", original_function=litellm.responses, diff --git a/tests/llm_translation/base_audio_transcription_unit_tests.py b/tests/llm_translation/base_audio_transcription_unit_tests.py index 76401b456fa..0234c05f853 100644 --- a/tests/llm_translation/base_audio_transcription_unit_tests.py +++ b/tests/llm_translation/base_audio_transcription_unit_tests.py @@ -55,7 +55,7 @@ class BaseLLMAudioTranscriptionTest(ABC): """ litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() AUDIO_FILE = open(file_path, "rb") transcription_call_args = self.get_base_audio_transcription_call_args() transcript = await litellm.atranscription( diff --git a/tests/llm_translation/base_llm_unit_tests.py b/tests/llm_translation/base_llm_unit_tests.py index 1a33422a31c..aadb093c050 100644 --- a/tests/llm_translation/base_llm_unit_tests.py +++ b/tests/llm_translation/base_llm_unit_tests.py @@ -144,7 +144,7 @@ class BaseLLMChatTest(ABC): assert response.choices[0].message.content is not None def test_tool_call_with_property_type_array(self): - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.utils import supports_function_calling os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" @@ -190,7 +190,7 @@ class BaseLLMChatTest(ABC): @pytest.mark.flaky(retries=3, delay=1) def test_tool_call_with_empty_enum_property(self): - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.utils import supports_function_calling os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" @@ -301,7 +301,7 @@ class BaseLLMChatTest(ABC): os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") - litellm._turn_on_debug() + litellm.turn_on_debug() base_completion_call_args = self.get_base_completion_call_args() @@ -327,7 +327,7 @@ class BaseLLMChatTest(ABC): os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") - litellm._turn_on_debug() + litellm.turn_on_debug() base_completion_call_args = self.get_base_completion_call_args() @@ -356,7 +356,7 @@ class BaseLLMChatTest(ABC): os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") - litellm._turn_on_debug() + litellm.turn_on_debug() image_content = [ {"type": "text", "text": "What's this file about?"}, @@ -395,7 +395,7 @@ class BaseLLMChatTest(ABC): os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") - litellm._turn_on_debug() + litellm.turn_on_debug() image_content = [ {"type": "text", "text": "What's this file about?"}, @@ -602,7 +602,7 @@ class BaseLLMChatTest(ABC): @pytest.mark.flaky(retries=6, delay=1) def test_json_response_pydantic_obj(self): - litellm._turn_on_debug() + litellm.turn_on_debug() from pydantic import BaseModel from litellm.utils import supports_response_schema @@ -698,7 +698,7 @@ class BaseLLMChatTest(ABC): """ PROD Test: ensure nested json schema sent to proxy works as expected. """ - litellm._turn_on_debug() + litellm.turn_on_debug() from pydantic import BaseModel from litellm.utils import supports_response_schema from litellm.llms.base_llm.base_utils import type_to_response_format_param @@ -751,7 +751,7 @@ class BaseLLMChatTest(ABC): """ from litellm.utils import supports_audio_input - litellm._turn_on_debug() + litellm.turn_on_debug() base_completion_call_args = self.get_base_completion_call_args() if not supports_audio_input(base_completion_call_args["model"], None): pytest.skip( @@ -1090,7 +1090,7 @@ class BaseLLMChatTest(ABC): from litellm import completion, ModelResponse litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.utils import supports_function_calling os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" @@ -1123,7 +1123,7 @@ class BaseLLMChatTest(ABC): from litellm import completion, ModelResponse litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.utils import supports_function_calling os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" @@ -1245,7 +1245,7 @@ class BaseLLMChatTest(ABC): async def test_completion_cost(self): from litellm import completion_cost - litellm._turn_on_debug() + litellm.turn_on_debug() os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") @@ -1332,7 +1332,7 @@ class BaseLLMChatTest(ABC): from litellm.utils import supports_function_calling from litellm import completion - litellm._turn_on_debug() + litellm.turn_on_debug() try: os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" @@ -1462,7 +1462,7 @@ class BaseLLMChatTest(ABC): ) in json.dumps(optional_params) try: - litellm._turn_on_debug() + litellm.turn_on_debug() response = completion( **base_completion_call_args, reasoning_effort="low", @@ -1665,7 +1665,7 @@ class BaseAnthropicChatTest(ABC): def test_completion_thinking_with_response_format(self): from pydantic import BaseModel - litellm._turn_on_debug() + litellm.turn_on_debug() class RFormat(BaseModel): question: str @@ -1685,7 +1685,7 @@ class BaseAnthropicChatTest(ABC): def test_completion_thinking_with_max_tokens(self): from pydantic import BaseModel - litellm._turn_on_debug() + litellm.turn_on_debug() base_completion_call_args = self.get_base_completion_call_args_with_thinking() @@ -1701,7 +1701,7 @@ class BaseAnthropicChatTest(ABC): def test_completion_thinking_without_max_tokens(self): from pydantic import BaseModel - litellm._turn_on_debug() + litellm.turn_on_debug() base_completion_call_args = self.get_base_completion_call_args_with_thinking() @@ -1714,7 +1714,7 @@ class BaseAnthropicChatTest(ABC): print(response) def test_completion_with_thinking_basic(self): - litellm._turn_on_debug() + litellm.turn_on_debug() base_completion_call_args = self.get_base_completion_call_args_with_thinking() messages = [{"role": "user", "content": "Generate 5 question + answer pairs"}] @@ -1815,7 +1815,7 @@ class BaseReasoningLLMTests(ABC): - Assert that `reasoning_content` is not None from response message - Assert that `reasoning_tokens` is greater than 0 from usage """ - litellm._turn_on_debug() + litellm.turn_on_debug() base_completion_call_args = self.get_base_completion_call_args() response: ModelResponse = self.completion_function( **base_completion_call_args, reasoning_effort="low" @@ -1835,7 +1835,7 @@ class BaseReasoningLLMTests(ABC): - Assert that `reasoning_content` is not None from streaming response - Assert that `reasoning_tokens` is greater than 0 from usage """ - # litellm._turn_on_debug() + # litellm.turn_on_debug() base_completion_call_args = self.get_base_completion_call_args() response: CustomStreamWrapper = self.completion_function( **base_completion_call_args, diff --git a/tests/llm_translation/base_rerank_unit_tests.py b/tests/llm_translation/base_rerank_unit_tests.py index df7dd33d7b0..e3c66dcc0be 100644 --- a/tests/llm_translation/base_rerank_unit_tests.py +++ b/tests/llm_translation/base_rerank_unit_tests.py @@ -90,7 +90,7 @@ class BaseLLMRerankTest(ABC): @pytest.mark.asyncio() @pytest.mark.parametrize("sync_mode", [True, False]) async def test_basic_rerank(self, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") rerank_call_args = self.get_base_rerank_call_args() diff --git a/tests/llm_translation/realtime/base_realtime_tests.py b/tests/llm_translation/realtime/base_realtime_tests.py index dabc66bb383..fd04dcb035f 100644 --- a/tests/llm_translation/realtime/base_realtime_tests.py +++ b/tests/llm_translation/realtime/base_realtime_tests.py @@ -178,7 +178,7 @@ class BaseRealtimeTest(ABC): 2. Initial event is received 3. Messages are properly forwarded """ - litellm._turn_on_debug() + litellm.turn_on_debug() if self.should_skip(): pytest.skip(self.get_skip_reason()) @@ -241,7 +241,7 @@ class BaseRealtimeTest(ABC): Test realtime connection with explicit query parameters. Verifies that query params are properly passed to the backend. """ - litellm._turn_on_debug() + litellm.turn_on_debug() if self.should_skip(): pytest.skip(self.get_skip_reason()) @@ -303,7 +303,7 @@ class BaseRealtimeTest(ABC): if self.should_skip(): pytest.skip(self.get_skip_reason()) - litellm._turn_on_debug() + litellm.turn_on_debug() # Create a custom websocket client that sends a message class InteractiveWebSocketClient(RealTimeWebSocketClient): diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index 396cd74f75a..a79240d477d 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -356,7 +356,7 @@ def test_process_anthropic_headers_with_no_matching_headers(): def test_anthropic_tool_use(tool_type, tool_config, message_content): """Test Anthropic tool use with computer use and web fetch tools.""" - litellm._turn_on_debug() + litellm.turn_on_debug() tools = [tool_config] model = "claude-sonnet-4-5-20250929" @@ -1015,7 +1015,7 @@ def test_anthropic_citations_api_streaming(): ) def test_anthropic_thinking_output(model): - litellm._turn_on_debug() + litellm.turn_on_debug() resp = completion( model=model, @@ -1045,7 +1045,7 @@ def test_anthropic_thinking_output(model): def test_anthropic_thinking_output_stream(model): litellm.set_verbose = True try: - # litellm._turn_on_debug() + # litellm.turn_on_debug() resp = litellm.completion( model=model, messages=[{"role": "user", "content": "Tell me a joke."}], @@ -1183,7 +1183,7 @@ async def test_anthropic_api_max_completion_tokens(model: str): ], ) def test_anthropic_websearch(optional_params: dict): - litellm._turn_on_debug() + litellm.turn_on_debug() params = { "model": "anthropic/claude-sonnet-4-5-20250929", "messages": [ @@ -1209,7 +1209,7 @@ def test_anthropic_websearch(optional_params: dict): def test_anthropic_text_editor(): - litellm._turn_on_debug() + litellm.turn_on_debug() params = { "model": "anthropic/claude-sonnet-4-5-20250929", "messages": [ @@ -1236,7 +1236,7 @@ def test_anthropic_text_editor(): os.getenv("ZAPIER_CI_CD_MCP_TOKEN") is None, reason="ZAPIER_CI_CD_MCP_TOKEN not set" ) def test_anthropic_mcp_server_tool_use(spec: str): - litellm._turn_on_debug() + litellm.turn_on_debug() if spec == "anthropic": tools = [ @@ -1282,7 +1282,7 @@ def test_anthropic_mcp_server_tool_use(spec: str): def test_anthropic_mcp_server_responses_api(model: str): from litellm import responses - litellm._turn_on_debug() + litellm.turn_on_debug() tools = [ { "type": "mcp", diff --git a/tests/llm_translation/test_azure_ai.py b/tests/llm_translation/test_azure_ai.py index f00409f280b..e49930dee1a 100644 --- a/tests/llm_translation/test_azure_ai.py +++ b/tests/llm_translation/test_azure_ai.py @@ -248,7 +248,7 @@ async def test_azure_ai_request_format(): """ from openai import AsyncAzureOpenAI, AzureOpenAI - litellm._turn_on_debug() + litellm.turn_on_debug() # Set up the test parameters api_key = os.getenv("AZURE_AI_API_KEY") @@ -272,7 +272,7 @@ async def test_azure_ai_request_format(): @pytest.mark.asyncio @pytest.mark.parametrize("model", ["azure/gpt5_series/gpt-5-mini"]) async def test_azure_gpt5_reasoning(model): - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.acompletion( model=model, messages=[{"role": "user", "content": "What is the capital of France?"}], @@ -352,7 +352,7 @@ async def test_azure_ai_model_router(): calculate_azure_model_router_flat_cost, ) - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.acompletion( model="azure_ai/model_router/azure-model-router", messages=[{"role": "user", "content": "hi who is this"}], @@ -391,7 +391,7 @@ async def test_azure_ai_model_router_streaming_model_in_chunk(): Test that Azure AI model router streaming returns the actual model in each chunk. The response should contain the actual model used (e.g., gpt-4.1-nano) not the request model (azure-model-router). """ - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.acompletion( model="azure_ai/azure-model-router", messages=[{"role": "user", "content": "hi"}], @@ -469,7 +469,7 @@ async def test_azure_ai_model_router_streaming_cost_with_stream_options(): litellm.callbacks = [test_callback] try: - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.acompletion( model="azure_ai/azure-model-router", messages=[{"role": "user", "content": "hi"}], diff --git a/tests/llm_translation/test_azure_o_series.py b/tests/llm_translation/test_azure_o_series.py index 2ee9bdb2be2..1c856df4b19 100644 --- a/tests/llm_translation/test_azure_o_series.py +++ b/tests/llm_translation/test_azure_o_series.py @@ -184,7 +184,7 @@ async def test_azure_o1_series_response_format_extra_params(): """ Tool calling should work for all azure o_series models. """ - litellm._turn_on_debug() + litellm.turn_on_debug() from openai import AsyncAzureOpenAI diff --git a/tests/llm_translation/test_azure_openai.py b/tests/llm_translation/test_azure_openai.py index df1892638b0..a6c005a0be7 100644 --- a/tests/llm_translation/test_azure_openai.py +++ b/tests/llm_translation/test_azure_openai.py @@ -613,7 +613,7 @@ def test_azure_safety_result(): """Bubble up safety result from Azure OpenAI""" from litellm import completion - litellm._turn_on_debug() + litellm.turn_on_debug() response = completion( model="azure/gpt-4.1-mini", @@ -631,7 +631,7 @@ def test_azure_openai_responses_bridge(): from litellm import completion import litellm - litellm._turn_on_debug() + litellm.turn_on_debug() with patch.object(litellm, "responses") as mock_responses: try: diff --git a/tests/llm_translation/test_bedrock_agentcore.py b/tests/llm_translation/test_bedrock_agentcore.py index 0087eb5b326..4b0a47d7e70 100644 --- a/tests/llm_translation/test_bedrock_agentcore.py +++ b/tests/llm_translation/test_bedrock_agentcore.py @@ -24,7 +24,7 @@ def test_bedrock_agentcore_basic(model): """ Test AgentCore invocation parameterized by model """ - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.completion( model=model, messages=[ @@ -49,7 +49,7 @@ async def test_bedrock_agentcore_with_streaming(model): Test AgentCore with streaming """ print("running streming test for model=", model) - # litellm._turn_on_debug() + # litellm.turn_on_debug() response = await litellm.acompletion( model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC", messages=[ @@ -71,7 +71,7 @@ def test_bedrock_agentcore_with_custom_params(): """ import json - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.llms.custom_httpx.http_handler import HTTPHandler client = HTTPHandler() @@ -139,7 +139,7 @@ def test_bedrock_agentcore_with_runtime_user_id(): """ import json - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.llms.custom_httpx.http_handler import HTTPHandler client = HTTPHandler() @@ -178,7 +178,7 @@ def test_bedrock_agentcore_with_session_and_user(): """ import json - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.llms.custom_httpx.http_handler import HTTPHandler client = HTTPHandler() @@ -222,7 +222,7 @@ def test_bedrock_agentcore_with_api_key_bearer_token(): """ import json - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.llms.custom_httpx.http_handler import HTTPHandler client = HTTPHandler() @@ -270,7 +270,7 @@ def test_bedrock_agentcore_with_all_parameters(): """ import json - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.llms.custom_httpx.http_handler import HTTPHandler client = HTTPHandler() @@ -339,7 +339,7 @@ def test_bedrock_agentcore_without_api_key_uses_sigv4(): """ import json - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.llms.custom_httpx.http_handler import HTTPHandler client = HTTPHandler() @@ -602,7 +602,7 @@ def test_agentcore_synchronous_non_streaming_response(): """ from litellm.llms.custom_httpx.http_handler import HTTPHandler - litellm._turn_on_debug() + litellm.turn_on_debug() client = HTTPHandler() # Mock a JSON response (typical for synchronous AgentCore calls) diff --git a/tests/llm_translation/test_bedrock_agents.py b/tests/llm_translation/test_bedrock_agents.py index 1685dd220d2..43b9237e0c6 100644 --- a/tests/llm_translation/test_bedrock_agents.py +++ b/tests/llm_translation/test_bedrock_agents.py @@ -16,7 +16,7 @@ import pytest @pytest.mark.asyncio @pytest.mark.skip(reason="Skipping bedrock agents test - arn not working") async def test_bedrock_agents(): - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.completion( model="bedrock/agent/L1RT58GYRW/MFPSBCXYTW", messages=[{"role": "user", "content": "Hi just respond with a ping message"}], @@ -41,7 +41,7 @@ async def test_bedrock_agents(): @pytest.mark.asyncio @pytest.mark.skip(reason="Skipping bedrock agents test - arn not working") async def test_bedrock_agents_with_streaming(): - # litellm._turn_on_debug() + # litellm.turn_on_debug() response = litellm.completion( model="bedrock/agent/L1RT58GYRW/MFPSBCXYTW", messages=[ @@ -60,7 +60,7 @@ async def test_bedrock_agents_with_streaming(): def test_bedrock_agents_with_custom_params(): - litellm._turn_on_debug() + litellm.turn_on_debug() from unittest.mock import MagicMock from litellm.llms.custom_httpx.http_handler import HTTPHandler diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 4161e08235b..b329a18ac72 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -505,7 +505,7 @@ def test_bedrock_system_prompt(system, model): def test_bedrock_claude_3_tool_calling(): try: litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() tools = [ { "type": "function", @@ -626,7 +626,7 @@ def test_completion_bedrock_mistral_completion_auth(): import os - litellm._turn_on_debug() + litellm.turn_on_debug() # aws_access_key_id = os.environ["AWS_ACCESS_KEY_ID"] # aws_secret_access_key = os.environ["AWS_SECRET_ACCESS_KEY"] @@ -2360,7 +2360,7 @@ async def test_bedrock_image_url_sync_client(): verbose_logger.setLevel(level=logging.DEBUG) - litellm._turn_on_debug() + litellm.turn_on_debug() client = AsyncHTTPHandler() messages = [ @@ -2461,7 +2461,7 @@ def test_bedrock_custom_deepseek(): from litellm.llms.custom_httpx.http_handler import HTTPHandler import json - litellm._turn_on_debug() + litellm.turn_on_debug() client = HTTPHandler() with patch.object(client, "post") as mock_post: @@ -2688,7 +2688,7 @@ def test_bedrock_description_param(): ) @pytest.mark.asyncio async def test_bedrock_thinking_in_assistant_message(sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler if sync_mode: @@ -2994,7 +2994,7 @@ async def test_bedrock_passthrough_router(): import litellm from litellm import Router - litellm._turn_on_debug() + litellm.turn_on_debug() router = Router( model_list=[ @@ -3054,7 +3054,7 @@ async def test_bedrock_converse__streaming_passthrough(monkeypatch): mock_custom_logger = MockCustomLogger() monkeypatch.setattr(litellm, "callbacks", [mock_custom_logger]) - litellm._turn_on_debug() + litellm.turn_on_debug() data = { "messages": [ @@ -3104,7 +3104,7 @@ async def test_bedrock_streaming_passthrough_test2(monkeypatch): mock_custom_logger = MockCustomLogger() monkeypatch.setattr(litellm, "callbacks", [mock_custom_logger]) - litellm._turn_on_debug() + litellm.turn_on_debug() data = { "max_tokens": 512, diff --git a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py b/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py index 184a0ae0749..8e473fdd110 100644 --- a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py +++ b/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py @@ -15,7 +15,7 @@ from litellm.llms.bedrock.common_utils import BedrockModelInfo def test_bedrock_completion_with_region_name(): - litellm._turn_on_debug() + litellm.turn_on_debug() client = HTTPHandler() with patch.object(client, "post") as mock_post: @@ -71,7 +71,7 @@ def test_bedrock_completion_with_region_name(): def test_bedrock_completion_with_dynamic_authentication_params(): - litellm._turn_on_debug() + litellm.turn_on_debug() client = HTTPHandler() with patch.object(client, "post") as mock_post: @@ -119,7 +119,7 @@ def test_bedrock_completion_with_dynamic_authentication_params(): def test_bedrock_completion_with_dynamic_bedrock_runtime_endpoint(): - litellm._turn_on_debug() + litellm.turn_on_debug() client = HTTPHandler() with patch.object(client, "post") as mock_post: diff --git a/tests/llm_translation/test_bedrock_embedding.py b/tests/llm_translation/test_bedrock_embedding.py index 1fc05b43b23..22b096a7a8c 100644 --- a/tests/llm_translation/test_bedrock_embedding.py +++ b/tests/llm_translation/test_bedrock_embedding.py @@ -106,7 +106,7 @@ def test_e2e_bedrock_embedding(): os.environ["AWS_REGION_NAME"] = "us-east-1" - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.embedding( model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0", input=["Hello world from LiteLLM with TwelveLabs Marengo!"], @@ -160,7 +160,7 @@ def test_e2e_bedrock_embedding_image_twelvelabs_marengo(): print("Testing image embedding...") original_region_name = os.environ.get("AWS_REGION_NAME") os.environ["AWS_REGION_NAME"] = "us-east-1" - litellm._turn_on_debug() + litellm.turn_on_debug() # Load duck.png and convert to base64 duck_img_path = os.path.join(os.path.dirname(__file__), "duck.png") @@ -229,7 +229,7 @@ def test_e2e_bedrock_async_invoke_embedding_twelvelabs_marengo(): print("Testing async invoke embedding...") original_region_name = os.environ.get("AWS_REGION_NAME") os.environ["AWS_REGION_NAME"] = "us-east-1" - litellm._turn_on_debug() + litellm.turn_on_debug() # Mock the HTTP call to return async invoke response with patch( @@ -292,7 +292,7 @@ async def test_e2e_bedrock_async_invoke_embedding_async_twelvelabs_marengo(): print("Testing async invoke embedding with async calls...") original_region_name = os.environ.get("AWS_REGION_NAME") os.environ["AWS_REGION_NAME"] = "us-east-1" - litellm._turn_on_debug() + litellm.turn_on_debug() # Mock the async HTTP call to return async invoke response with patch( diff --git a/tests/llm_translation/test_bedrock_invoke_tests.py b/tests/llm_translation/test_bedrock_invoke_tests.py index 46386b207cb..3f84723ed2f 100644 --- a/tests/llm_translation/test_bedrock_invoke_tests.py +++ b/tests/llm_translation/test_bedrock_invoke_tests.py @@ -19,7 +19,7 @@ _AWSMP_LOGO_IMAGE_URL = ( @pytest.mark.flaky(retries=3, delay=5) class TestBedrockInvokeClaudeJson(BaseLLMChatTest): def get_base_completion_call_args(self) -> dict: - litellm._turn_on_debug() + litellm.turn_on_debug() return { "model": "bedrock/invoke/us.anthropic.claude-haiku-4-5-20251001-v1:0", } diff --git a/tests/llm_translation/test_bedrock_llama.py b/tests/llm_translation/test_bedrock_llama.py index b02b482b955..9fbf108fbc4 100644 --- a/tests/llm_translation/test_bedrock_llama.py +++ b/tests/llm_translation/test_bedrock_llama.py @@ -13,7 +13,7 @@ class TestBedrockTestSuite(BaseLLMChatTest): pass def get_base_completion_call_args(self) -> dict: - litellm._turn_on_debug() + litellm.turn_on_debug() return { "model": "bedrock/converse/us.meta.llama3-3-70b-instruct-v1:0", } diff --git a/tests/llm_translation/test_bedrock_moonshot.py b/tests/llm_translation/test_bedrock_moonshot.py index 5323a87c366..bb0f3510e68 100644 --- a/tests/llm_translation/test_bedrock_moonshot.py +++ b/tests/llm_translation/test_bedrock_moonshot.py @@ -33,7 +33,7 @@ class TestBedrockMoonshotInvoke(BaseLLMChatTest): test_json_response_format_stream = None def get_base_completion_call_args(self) -> dict: - litellm._turn_on_debug() + litellm.turn_on_debug() return { "model": "bedrock/invoke/moonshot.kimi-k2-thinking", } diff --git a/tests/llm_translation/test_bedrock_nova_json.py b/tests/llm_translation/test_bedrock_nova_json.py index f9531c99b52..be8fd321316 100644 --- a/tests/llm_translation/test_bedrock_nova_json.py +++ b/tests/llm_translation/test_bedrock_nova_json.py @@ -14,7 +14,7 @@ class TestBedrockNovaJson(BaseLLMChatTest): test_tool_call_with_property_type_array = None def get_base_completion_call_args(self) -> dict: - litellm._turn_on_debug() + litellm.turn_on_debug() return { "model": "bedrock/converse/us.amazon.nova-micro-v1:0", } diff --git a/tests/llm_translation/test_cohere.py b/tests/llm_translation/test_cohere.py index 729f42f8984..3ee3ab6ad9d 100644 --- a/tests/llm_translation/test_cohere.py +++ b/tests/llm_translation/test_cohere.py @@ -267,7 +267,7 @@ async def test_cohere_request_body_with_allowed_params(): def test_cohere_embedding_outout_dimensions(): - litellm._turn_on_debug() + litellm.turn_on_debug() response = embedding( model="cohere/embed-v4.0", input="Hello, world!", dimensions=512 ) diff --git a/tests/llm_translation/test_deepseek_completion.py b/tests/llm_translation/test_deepseek_completion.py index 2ede5d3f3f8..5838be8fc85 100644 --- a/tests/llm_translation/test_deepseek_completion.py +++ b/tests/llm_translation/test_deepseek_completion.py @@ -24,7 +24,7 @@ def test_deepseek_mock_completion(stream): import litellm from litellm import completion - litellm._turn_on_debug() + litellm.turn_on_debug() response = completion( model="deepseek/deepseek-reasoner", @@ -52,7 +52,7 @@ async def test_deepseek_provider_async_completion(stream): from unittest.mock import patch, AsyncMock, MagicMock from litellm import acompletion - litellm._turn_on_debug() + litellm.turn_on_debug() # Set up the test parameters api_key = "fake_api_key" diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index 7b0b741563d..6154df547ae 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -106,7 +106,7 @@ class TestGoogleAIStudioGemini(BaseLLMChatTest): os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") - litellm._turn_on_debug() + litellm.turn_on_debug() base_completion_call_args = self.get_base_completion_call_args() @@ -335,7 +335,7 @@ def test_gemini_context_caching_separate_messages(): def test_gemini_image_generation(): - # litellm._turn_on_debug() + # litellm.turn_on_debug() response = completion( model="gemini/gemini-2.5-flash-image", messages=[{"role": "user", "content": "Generate an image of a cat"}], @@ -620,7 +620,7 @@ def test_gemini_imagen_models_use_predict_endpoint(): def test_gemini_thinking(): - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.types.utils import Message, CallTypes from litellm.utils import return_raw_request import json @@ -660,7 +660,7 @@ def test_gemini_thinking(): def test_gemini_thinking_budget_0(): - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.types.utils import Message, CallTypes from litellm.utils import return_raw_request import json @@ -686,7 +686,7 @@ def test_gemini_finish_reason(): import os from litellm import completion - litellm._turn_on_debug() + litellm.turn_on_debug() response = completion( model="gemini/gemini-2.5-flash-lite", messages=[{"role": "user", "content": "give me 3 random words"}], @@ -701,7 +701,7 @@ def test_gemini_finish_reason(): def test_gemini_url_context(): from litellm import completion - litellm._turn_on_debug() + litellm.turn_on_debug() URL1 = "https://www.foodnetwork.com/recipes/ina-garten/perfect-roast-chicken-recipe-1940592" prompt = f""" @@ -727,7 +727,7 @@ def test_gemini_url_context(): def test_gemini_with_grounding(): from litellm import completion, Usage, stream_chunk_builder - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True tools = [{"googleSearch": {}}] @@ -763,7 +763,7 @@ def test_gemini_with_grounding(): def test_gemini_with_empty_function_call_arguments(): from litellm import completion - litellm._turn_on_debug() + litellm.turn_on_debug() tools = [ { "type": "function", @@ -1026,7 +1026,7 @@ def test_gemini_tool_use(): @pytest.mark.asyncio async def test_gemini_image_generation_async(): - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.acompletion( messages=[ { @@ -1059,7 +1059,7 @@ async def test_gemini_image_generation_async(): @pytest.mark.asyncio async def test_gemini_image_generation_async_stream(): - # litellm._turn_on_debug() + # litellm.turn_on_debug() response = await litellm.acompletion( messages=[ { @@ -1129,7 +1129,7 @@ def get_current_weather(location, unit="fahrenheit"): def test_gemini_with_thinking(): from litellm import completion - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.modify_params = True model = "gemini/gemini-2.5-flash" messages = [ @@ -1429,7 +1429,7 @@ def l(status_code, expected_exception): def test_gemini_embedding(): - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.embedding( model="gemini/gemini-embedding-001", input="Hello, world!", diff --git a/tests/llm_translation/test_gpt4o_audio.py b/tests/llm_translation/test_gpt4o_audio.py index 0f20119e4ef..054af882a81 100644 --- a/tests/llm_translation/test_gpt4o_audio.py +++ b/tests/llm_translation/test_gpt4o_audio.py @@ -92,7 +92,7 @@ async def test_audio_input_to_model(stream, model): audio_format = "pcm16" if stream is False: audio_format = "wav" - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.drop_params = True url = "https://openaiassets.blob.core.windows.net/$web/API/docs/audio/alloy.wav" response = requests.get(url) diff --git a/tests/llm_translation/test_litellm_proxy_provider.py b/tests/llm_translation/test_litellm_proxy_provider.py index 8a010467889..ffc6f5b1180 100644 --- a/tests/llm_translation/test_litellm_proxy_provider.py +++ b/tests/llm_translation/test_litellm_proxy_provider.py @@ -128,7 +128,7 @@ async def _gateway_embedding_via_injected_client( @pytest.mark.asyncio async def test_litellm_gateway_from_sdk_embedding(is_async: bool): litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() transport, response = await _gateway_embedding_via_injected_client(is_async) @@ -152,7 +152,7 @@ async def test_litellm_gateway_from_sdk_embedding_under_foreign_cassette(tmp_pat @pytest.mark.parametrize("is_async", [False, True]) @pytest.mark.asyncio async def test_litellm_gateway_from_sdk_image_generation(is_async): - litellm._turn_on_debug() + litellm.turn_on_debug() if is_async: from openai import AsyncOpenAI @@ -202,7 +202,7 @@ async def test_litellm_gateway_from_sdk_image_generation(is_async): @pytest.mark.asyncio async def test_litellm_gateway_image_generation_direct(is_async): """Test image generation using the litellm_proxy provider directly.""" - litellm._turn_on_debug() + litellm.turn_on_debug() # Create mock response that matches OpenAI's response structure mock_openai_response = MagicMock() @@ -276,7 +276,7 @@ async def test_litellm_gateway_image_generation_direct(is_async): @pytest.mark.parametrize("is_async", [False, True]) @pytest.mark.asyncio async def test_litellm_gateway_from_sdk_image_edit(is_async): - litellm._turn_on_debug() + litellm.turn_on_debug() mock_response = { "created": 1, @@ -331,7 +331,7 @@ async def test_litellm_gateway_from_sdk_image_edit(is_async): @pytest.mark.asyncio async def test_litellm_gateway_from_sdk_transcription(is_async): litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() if is_async: from openai import AsyncOpenAI @@ -427,7 +427,7 @@ async def test_litellm_gateway_from_sdk_speech(is_async): @pytest.mark.asyncio async def test_litellm_gateway_from_sdk_rerank(is_async): litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() if is_async: client = AsyncHTTPHandler() @@ -521,7 +521,7 @@ async def test_litellm_gateway_from_sdk_rerank(is_async): def test_litellm_gateway_from_sdk_with_response_cost_in_additional_headers(): litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() from openai import OpenAI diff --git a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py index 31c554985a7..610fd08162a 100644 --- a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py +++ b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py @@ -1794,41 +1794,41 @@ class TestSafeConvertCreatedField: import time from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _safe_convert_created_field, + safe_convert_created_field, ) - result = _safe_convert_created_field(None) + result = safe_convert_created_field(None) assert abs(result - int(time.time())) <= 1 def test_int_passthrough(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _safe_convert_created_field, + safe_convert_created_field, ) - assert _safe_convert_created_field(1700000000) == 1700000000 + assert safe_convert_created_field(1700000000) == 1700000000 def test_float_truncated(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _safe_convert_created_field, + safe_convert_created_field, ) - assert _safe_convert_created_field(1700000000.999) == 1700000000 + assert safe_convert_created_field(1700000000.999) == 1700000000 def test_string_converted(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _safe_convert_created_field, + safe_convert_created_field, ) - assert _safe_convert_created_field("1700000000.5") == 1700000000 + assert safe_convert_created_field("1700000000.5") == 1700000000 def test_invalid_string_returns_current_time(self): import time from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _safe_convert_created_field, + safe_convert_created_field, ) - result = _safe_convert_created_field("not-a-number") + result = safe_convert_created_field("not-a-number") assert abs(result - int(time.time())) <= 1 @@ -2003,14 +2003,14 @@ class TestConvertToStreamingResponseAsync: class TestHandleInvalidParallelToolCalls: def test_none_input(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _handle_invalid_parallel_tool_calls, + handle_invalid_parallel_tool_calls, ) - assert _handle_invalid_parallel_tool_calls(None) is None + assert handle_invalid_parallel_tool_calls(None) is None def test_normal_tool_calls_unchanged(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _handle_invalid_parallel_tool_calls, + handle_invalid_parallel_tool_calls, ) from litellm.types.utils import ChatCompletionMessageToolCall, Function @@ -2021,13 +2021,13 @@ class TestHandleInvalidParallelToolCalls: function=Function(name="get_weather", arguments='{"city": "NYC"}'), ) ] - result = _handle_invalid_parallel_tool_calls(tool_calls) + result = handle_invalid_parallel_tool_calls(tool_calls) assert len(result) == 1 assert result[0].function.name == "get_weather" def test_multi_tool_use_parallel_expanded(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _handle_invalid_parallel_tool_calls, + handle_invalid_parallel_tool_calls, ) from litellm.types.utils import ChatCompletionMessageToolCall, Function @@ -2054,7 +2054,7 @@ class TestHandleInvalidParallelToolCalls: ), ) ] - result = _handle_invalid_parallel_tool_calls(tool_calls) + result = handle_invalid_parallel_tool_calls(tool_calls) assert len(result) == 2 assert result[0].function.name == "get_weather" assert result[0].id == "call_1_0" @@ -2064,7 +2064,7 @@ class TestHandleInvalidParallelToolCalls: def test_invalid_json_returns_original(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _handle_invalid_parallel_tool_calls, + handle_invalid_parallel_tool_calls, ) from litellm.types.utils import ChatCompletionMessageToolCall, Function @@ -2075,7 +2075,7 @@ class TestHandleInvalidParallelToolCalls: function=Function(name="some_func", arguments="not valid json{{{"), ) ] - result = _handle_invalid_parallel_tool_calls(tool_calls) + result = handle_invalid_parallel_tool_calls(tool_calls) assert len(result) == 1 assert result[0].id == "call_1" @@ -2083,13 +2083,13 @@ class TestHandleInvalidParallelToolCalls: class TestShouldConvertToolCallToJsonMode: def test_returns_true_when_conditions_met(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _should_convert_tool_call_to_json_mode, + should_convert_tool_call_to_json_mode, ) from litellm.constants import RESPONSE_FORMAT_TOOL_NAME tool_calls = [{"function": {"name": RESPONSE_FORMAT_TOOL_NAME}}] assert ( - _should_convert_tool_call_to_json_mode( + should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=True ) is True @@ -2097,13 +2097,13 @@ class TestShouldConvertToolCallToJsonMode: def test_returns_false_when_flag_off(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _should_convert_tool_call_to_json_mode, + should_convert_tool_call_to_json_mode, ) from litellm.constants import RESPONSE_FORMAT_TOOL_NAME tool_calls = [{"function": {"name": RESPONSE_FORMAT_TOOL_NAME}}] assert ( - _should_convert_tool_call_to_json_mode( + should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=False ) is False @@ -2111,12 +2111,12 @@ class TestShouldConvertToolCallToJsonMode: def test_returns_false_when_wrong_tool_name(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _should_convert_tool_call_to_json_mode, + should_convert_tool_call_to_json_mode, ) tool_calls = [{"function": {"name": "some_other_tool"}}] assert ( - _should_convert_tool_call_to_json_mode( + should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=True ) is False @@ -2124,7 +2124,7 @@ class TestShouldConvertToolCallToJsonMode: def test_returns_false_when_multiple_tool_calls(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _should_convert_tool_call_to_json_mode, + should_convert_tool_call_to_json_mode, ) from litellm.constants import RESPONSE_FORMAT_TOOL_NAME @@ -2133,7 +2133,7 @@ class TestShouldConvertToolCallToJsonMode: {"function": {"name": "other"}}, ] assert ( - _should_convert_tool_call_to_json_mode( + should_convert_tool_call_to_json_mode( tool_calls=tool_calls, convert_tool_call_to_json_mode=True ) is False @@ -2141,11 +2141,11 @@ class TestShouldConvertToolCallToJsonMode: def test_returns_false_when_none(self): from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _should_convert_tool_call_to_json_mode, + should_convert_tool_call_to_json_mode, ) assert ( - _should_convert_tool_call_to_json_mode( + should_convert_tool_call_to_json_mode( tool_calls=None, convert_tool_call_to_json_mode=True ) is False diff --git a/tests/llm_translation/test_openai.py b/tests/llm_translation/test_openai.py index d748a56e90c..f345717802d 100644 --- a/tests/llm_translation/test_openai.py +++ b/tests/llm_translation/test_openai.py @@ -433,7 +433,7 @@ def validate_web_search_annotations(annotations: ChatCompletionAnnotation): @pytest.mark.flaky(reruns=3) def test_openai_web_search(): """Makes a simple web search request and validates the response contains web search annotations and all expected fields are present""" - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.completion( model="openai/gpt-5-search-api", messages=[ @@ -452,7 +452,7 @@ def test_openai_web_search(): def test_openai_web_search_streaming(): """Makes a simple web search request and validates the response contains web search annotations and all expected fields are present""" - # litellm._turn_on_debug() + # litellm.turn_on_debug() test_openai_web_search: Optional[ChatCompletionAnnotation] = None response = litellm.completion( model="openai/gpt-5-search-api", @@ -624,7 +624,7 @@ def test_openai_responses_only_model_bridge(): """ Test that the responses-only model bridge works correctly """ - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.completion( model="gpt-5.5-pro", messages=[{"role": "user", "content": "Hey, how's it going?"}], @@ -811,7 +811,7 @@ def test_openai_service_tier_parameter_sync(): def test_gpt_5_reasoning_streaming(): - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.completion( model="openai/responses/gpt-5-mini", messages=[{"role": "user", "content": "Think of a poem, and then write it."}], @@ -832,7 +832,7 @@ def test_gpt_5_reasoning_streaming(): def test_openai_gpt_5_codex_reasoning(): - litellm._turn_on_debug() + litellm.turn_on_debug() completion_kwargs = { "model": "gpt-5.3-codex", "messages": [ diff --git a/tests/llm_translation/test_openrouter.py b/tests/llm_translation/test_openrouter.py index 8ecf9b4a8a2..5bdf3316f2e 100644 --- a/tests/llm_translation/test_openrouter.py +++ b/tests/llm_translation/test_openrouter.py @@ -4,7 +4,7 @@ import litellm def test_completion_openrouter_reasoning_content(): - litellm._turn_on_debug() + litellm.turn_on_debug() resp = litellm.completion( model="openrouter/anthropic/claude-sonnet-4", messages=[{"role": "user", "content": "Hello world"}], @@ -15,7 +15,7 @@ def test_completion_openrouter_reasoning_content(): def test_completion_openrouter_image_generation(): - litellm._turn_on_debug() + litellm.turn_on_debug() resp = litellm.completion( model="openrouter/google/gemini-2.5-flash-image", messages=[{"role": "user", "content": "Generate an image of a cat"}], @@ -29,7 +29,7 @@ def test_completion_openrouter_image_generation(): def test_openrouter_embedding(): """Test OpenRouter embeddings support.""" - litellm._turn_on_debug() + litellm.turn_on_debug() resp = litellm.embedding( model="openrouter/openai/text-embedding-3-small", input=["Hello world", "How are you?"], diff --git a/tests/llm_translation/test_optional_params.py b/tests/llm_translation/test_optional_params.py index 58446014bdf..7cf2210572a 100644 --- a/tests/llm_translation/test_optional_params.py +++ b/tests/llm_translation/test_optional_params.py @@ -158,13 +158,13 @@ def test_allowed_openai_params_does_not_forward_unset_params(): added ``optional_params["enable_thinking"] = None`` which then crashed the openai client. """ - from litellm.utils import _apply_openai_param_overrides + from litellm.utils import apply_openai_param_overrides chat_template_kwargs = {"enable_thinking": False} optional_params: dict = {} non_default_params = {"chat_template_kwargs": chat_template_kwargs} - result = _apply_openai_param_overrides( + result = apply_openai_param_overrides( optional_params=optional_params, non_default_params=non_default_params, allowed_openai_params=["chat_template_kwargs", "enable_thinking"], diff --git a/tests/llm_translation/test_skills_api.py b/tests/llm_translation/test_skills_api.py index d21e7376ea7..64e34e2a7e9 100644 --- a/tests/llm_translation/test_skills_api.py +++ b/tests/llm_translation/test_skills_api.py @@ -111,7 +111,7 @@ class BaseSkillsAPITest(ABC): pytest.skip(f"No API key provided for {custom_llm_provider}") litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() # Use helper to create skill zip skill_name = "test-skill-litellm" diff --git a/tests/local_testing/cache_unit_tests.py b/tests/local_testing/cache_unit_tests.py index a1973d477b2..515d668896d 100644 --- a/tests/local_testing/cache_unit_tests.py +++ b/tests/local_testing/cache_unit_tests.py @@ -28,7 +28,7 @@ class LLMCachingUnitTests(ABC): @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_cache_completion(self, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() random_number = random.randint( 1, 100000 @@ -142,7 +142,7 @@ class LLMCachingUnitTests(ABC): @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_disk_cache_embedding(self, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() random_number = random.randint( 1, 100000 diff --git a/tests/local_testing/test_amazing_vertex_completion.py b/tests/local_testing/test_amazing_vertex_completion.py index 7f4044fc87e..5dc6b0c174c 100644 --- a/tests/local_testing/test_amazing_vertex_completion.py +++ b/tests/local_testing/test_amazing_vertex_completion.py @@ -322,7 +322,7 @@ def test_avertex_ai_stream(): async def test_async_vertexai_streaming_response(): import random - litellm._turn_on_debug() + litellm.turn_on_debug() load_vertex_ai_credentials() test_models = ( @@ -3301,7 +3301,7 @@ def test_vertex_ai_llama_tool_calling(): os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") load_vertex_ai_credentials() - litellm._turn_on_debug() + litellm.turn_on_debug() args = { "model": "vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas", "messages": [ @@ -3347,7 +3347,7 @@ def test_gemini_nullable_object_tool_schema_httpx(): Ensure nullable object tool params preserve nested properties in Vertex schema conversion. """ load_vertex_ai_credentials() - litellm._turn_on_debug() + litellm.turn_on_debug() tools = [ { @@ -3503,7 +3503,7 @@ def test_vertex_ai_streaming_response_id(): def test_vertex_ai_gemini_2_5_pro_streaming(): try: load_vertex_ai_credentials() - # litellm._turn_on_debug() + # litellm.turn_on_debug() response = completion( model="vertex_ai/gemini-2.5-pro", messages=[{"role": "user", "content": "Hi!"}], @@ -3613,7 +3613,7 @@ def test_vertex_ai_gemini_audio_ogg(): async def test_vertex_ai_deepseek(): """Test that deepseek models use the correct v1 API endpoint instead of v1beta1.""" load_vertex_ai_credentials() - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler client = AsyncHTTPHandler() @@ -3655,7 +3655,7 @@ def test_gemini_grounding_on_streaming(): from litellm import completion load_vertex_ai_credentials() - # litellm._turn_on_debug() + # litellm.turn_on_debug() args = { "model": "vertex_ai/gemini-3-flash-preview", "messages": [ @@ -3689,7 +3689,7 @@ def test_gemini_google_maps_tool_simple(): Test googleMaps tool with just enableWidget parameter. """ load_vertex_ai_credentials() - litellm._turn_on_debug() + litellm.turn_on_debug() tools = [{"googleMaps": {"enableWidget": True}}] tools_with_location = [ diff --git a/tests/local_testing/test_anthropic_prompt_caching.py b/tests/local_testing/test_anthropic_prompt_caching.py index 904b3ead92d..8775656e785 100644 --- a/tests/local_testing/test_anthropic_prompt_caching.py +++ b/tests/local_testing/test_anthropic_prompt_caching.py @@ -203,7 +203,7 @@ def anthropic_messages(): @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_anthropic_vertex_ai_prompt_caching(anthropic_messages, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() load_vertex_ai_credentials() diff --git a/tests/local_testing/test_blocked_user_list.py b/tests/local_testing/test_blocked_user_list.py index f7913c83662..4c530d27c4c 100644 --- a/tests/local_testing/test_blocked_user_list.py +++ b/tests/local_testing/test_blocked_user_list.py @@ -24,7 +24,7 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.enterprise.enterprise_hooks.blocked_user_list import ( - _ENTERPRISE_BlockedUserList, + ENTERPRISE_BlockedUserList, ) from litellm.proxy.management_endpoints.internal_user_endpoints import ( new_user, @@ -101,7 +101,7 @@ async def test_block_user_check(prisma_client): litellm.blocked_user_list = ["user_id_1"] - blocked_user_obj = _ENTERPRISE_BlockedUserList( + blocked_user_obj = ENTERPRISE_BlockedUserList( prisma_client=litellm.proxy.proxy_server.prisma_client ) diff --git a/tests/local_testing/test_caching.py b/tests/local_testing/test_caching.py index 3e96896f47f..57c14f77bb6 100644 --- a/tests/local_testing/test_caching.py +++ b/tests/local_testing/test_caching.py @@ -110,7 +110,7 @@ async def test_batch_get_cache_with_none_keys(sync_mode): """ from litellm.caching.caching import RedisCache - litellm._turn_on_debug() + litellm.turn_on_debug() redis_cache = RedisCache( host=os.environ.get("REDIS_HOST"), @@ -505,7 +505,7 @@ def test_embedding_caching(): @pytest.mark.asyncio async def test_embedding_caching_individual_items_and_then_list(): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.cache = Cache() text_to_embed = [ "hello", @@ -2592,7 +2592,7 @@ async def test_redis_increment_pipeline(): from litellm.caching.redis_cache import RedisCache litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() redis_cache = RedisCache( host=os.environ["REDIS_HOST"], port=os.environ["REDIS_PORT"], @@ -2863,7 +2863,7 @@ def test_caching_with_reasoning_content(): def test_caching_reasoning_args_miss(): # test in memory cache try: - # litellm._turn_on_debug() + # litellm.turn_on_debug() litellm.set_verbose = True litellm.cache = Cache() response1 = completion( @@ -2889,7 +2889,7 @@ def test_caching_reasoning_args_miss(): # test in memory cache def test_caching_reasoning_args_hit(): # test in memory cache try: - # litellm._turn_on_debug() + # litellm.turn_on_debug() litellm.set_verbose = True litellm.cache = Cache() response1 = completion( @@ -2916,7 +2916,7 @@ def test_caching_reasoning_args_hit(): # test in memory cache def test_caching_thinking_args_miss(): # test in memory cache try: - # litellm._turn_on_debug() + # litellm.turn_on_debug() litellm.set_verbose = True litellm.cache = Cache() response1 = completion( @@ -2942,7 +2942,7 @@ def test_caching_thinking_args_miss(): # test in memory cache def test_caching_thinking_args_hit(): # test in memory cache try: - # litellm._turn_on_debug() + # litellm.turn_on_debug() litellm.set_verbose = True litellm.cache = Cache() response1 = completion( diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py index 5ff7d79e3f8..fd0cda9920b 100644 --- a/tests/local_testing/test_completion.py +++ b/tests/local_testing/test_completion.py @@ -1011,7 +1011,7 @@ async def test_openai_compatible_custom_api_video(provider): def test_lm_studio_completion(monkeypatch): monkeypatch.delenv("LM_STUDIO_API_KEY", raising=False) monkeypatch.delenv("OPENAI_API_KEY", raising=False) - litellm._turn_on_debug() + litellm.turn_on_debug() try: completion( api_key="fake-key", @@ -1259,7 +1259,7 @@ def test_completion_openai(): @pytest.mark.flaky(retries=3, delay=1) def test_completion_openai_pydantic(model, api_version): try: - litellm._turn_on_debug() + litellm.turn_on_debug() from pydantic import BaseModel messages = [ @@ -3611,7 +3611,7 @@ def test_completion_novita_ai_dynamic_params(api_key): def test_deepseek_reasoning_content_completion(): try: litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() resp = litellm.completion( timeout=5, model="deepseek/deepseek-reasoner", @@ -3624,7 +3624,7 @@ def test_deepseek_reasoning_content_completion(): def test_qwen_text_completion(): - # litellm._turn_on_debug() + # litellm.turn_on_debug() resp = litellm.completion( model="text-completion-openai/gpt-5.4-nano", messages=[{"content": "hello", "role": "user"}], @@ -3692,7 +3692,7 @@ def test_completion_o3_mini_temperature(): def test_completion_gpt_4o_empty_str(): - litellm._turn_on_debug() + litellm.turn_on_debug() from openai import OpenAI from unittest.mock import MagicMock diff --git a/tests/local_testing/test_completion_cost.py b/tests/local_testing/test_completion_cost.py index c24e6c32369..d3838a3a264 100644 --- a/tests/local_testing/test_completion_cost.py +++ b/tests/local_testing/test_completion_cost.py @@ -978,7 +978,7 @@ def test_completion_cost_prompt_caching(model, custom_llm_provider): ) @pytest.mark.skip(reason="databricks is having an active outage") def test_completion_cost_databricks(model): - litellm._turn_on_debug() + litellm.turn_on_debug() os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") messages = [{"role": "user", "content": "What is 2+2?"}] @@ -2357,7 +2357,7 @@ def test_completion_cost_azure_tts(): def test_select_model_name_for_cost_calc(): - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc from litellm.types.utils import ModelResponse, Choices, Usage, Message args = { @@ -2393,7 +2393,7 @@ def test_select_model_name_for_cost_calc(): "custom_pricing": None, } - return_model = _select_model_name_for_cost_calc(**args) + return_model = select_model_name_for_cost_calc(**args) assert return_model == "azure_ai/mistral-large" @@ -2578,7 +2578,7 @@ def test_cost_calculator_with_base_model_with_router(base_model_arg): def test_cost_calculator_with_base_model_with_router_embedding(base_model_arg): from litellm import Router - litellm._turn_on_debug() + litellm.turn_on_debug() model_item = { "model_name": "random-model", diff --git a/tests/local_testing/test_docker_no_network_on_deploy.py b/tests/local_testing/test_docker_no_network_on_deploy.py index 8083eb912b4..b73b103eb7e 100644 --- a/tests/local_testing/test_docker_no_network_on_deploy.py +++ b/tests/local_testing/test_docker_no_network_on_deploy.py @@ -1,385 +1,385 @@ -""" -Test to ensure Docker container does not go out to network on deploy. - -This test verifies that the LiteLLM proxy container does not make outbound -network requests during startup. This is important for: -1. Air-gapped environments where outbound network is restricted -2. Security compliance requiring no unexpected network calls -3. Fast container startup without network dependencies - -The test works by: -1. Building/running the container with network disabled -2. Verifying the container starts successfully without network -3. Checking for any errors related to network failures during startup -""" - -import os -import re -import subprocess -import time - -import pytest - -from tests._master_key import MASTER_KEY - - -def is_docker_available() -> bool: - """Check if Docker is available and running.""" - try: - result = subprocess.run( - ["docker", "info"], - capture_output=True, - text=True, - timeout=10, - ) - return result.returncode == 0 - except (subprocess.TimeoutExpired, FileNotFoundError): - return False - - -@pytest.mark.skipif( - not is_docker_available(), - reason="Docker not available", -) -class TestDockerNoNetworkOnDeploy: - """ - Test suite for verifying Docker container starts without network access. - """ - - # Container and image names for testing - TEST_IMAGE_NAME = "litellm-no-network-test" - TEST_CONTAINER_NAME = "litellm-no-network-test-container" - - # Timeout for container operations - CONTAINER_START_TIMEOUT = 60 # seconds - - # Patterns that indicate network-related failures during startup - NETWORK_ERROR_PATTERNS = [ - r"connection refused", - r"network is unreachable", - r"could not resolve host", - r"name resolution failed", - r"dns lookup failed", - r"failed to establish.*connection", - r"socket.gaierror", - r"urllib.error.URLError.*Errno", - r"requests.exceptions.ConnectionError", - r"httpx.*ConnectError", - r"aiohttp.*ClientConnectorError", - ] - - # Patterns that indicate EXPECTED local-only startup - LOCAL_STARTUP_PATTERNS = [ - r"starting.*(server|proxy)", - r"listening on", - r"uvicorn running", - r"application startup complete", - ] - - @pytest.fixture(autouse=True) - def cleanup(self): - """Cleanup any existing test containers before and after each test.""" - self._cleanup_container() - yield - self._cleanup_container() - - def _cleanup_container(self): - """Remove test container if it exists.""" - subprocess.run( - ["docker", "rm", "-f", self.TEST_CONTAINER_NAME], - capture_output=True, - timeout=30, - ) - - def _build_test_image(self) -> bool: - """ - Build the Docker image if needed. - Returns True if image is available (built or already exists). - """ - # Check if main litellm image exists - result = subprocess.run( - ["docker", "images", "-q", "litellm/litellm"], - capture_output=True, - text=True, - timeout=10, - ) - if result.stdout.strip(): - # Image exists, use it - return True - - # Try to build from Dockerfile - dockerfile_path = os.path.join( - os.path.dirname(__file__), - "..", - "..", - "..", - "Dockerfile", - ) - if os.path.exists(dockerfile_path): - result = subprocess.run( - [ - "docker", - "build", - "-t", - self.TEST_IMAGE_NAME, - "-f", - dockerfile_path, - os.path.dirname(dockerfile_path), - ], - capture_output=True, - text=True, - timeout=600, # 10 minutes for build - ) - return result.returncode == 0 - - return False - - def test_container_starts_without_network(self): - """ - Test that the container can start with network completely disabled. - - This test runs the container with --network=none to ensure no outbound - network requests are required during startup. - """ - # Use a minimal config that doesn't require external services - minimal_config = f""" -model_list: - - model_name: fake-model - litellm_params: - model: fake/fake-model - -general_settings: - master_key: {MASTER_KEY} - database_url: null - -environment_variables: {{}} -""" - - # Create a temporary config file - config_path = "/tmp/litellm_test_config.yaml" - with open(config_path, "w") as f: - f.write(minimal_config) - - image_to_use = "litellm/litellm" - # Check if image exists - result = subprocess.run( - ["docker", "images", "-q", image_to_use], - capture_output=True, - text=True, - timeout=10, - ) - if not result.stdout.strip(): - pytest.skip(f"Docker image {image_to_use} not available") - - # Run container with network disabled - run_cmd = [ - "docker", - "run", - "--name", - self.TEST_CONTAINER_NAME, - "--network=none", # Disable all network access - "-v", - f"{config_path}:/app/config.yaml:ro", - "-e", - f"LITELLM_MASTER_KEY={MASTER_KEY}", - "-e", - "DATABASE_URL=", # Empty to disable DB - "-e", - "STORE_MODEL_IN_DB=false", - "-e", - "LITELLM_LOG=DEBUG", - "-d", # Detached mode - image_to_use, - "--config", - "/app/config.yaml", - ] - - result = subprocess.run( - run_cmd, - capture_output=True, - text=True, - timeout=30, - ) - - if result.returncode != 0: - pytest.fail(f"Failed to start container: {result.stderr}") - - # Wait for container to start up - time.sleep(5) - - # Check container logs for any network-related errors - logs_result = subprocess.run( - ["docker", "logs", self.TEST_CONTAINER_NAME], - capture_output=True, - text=True, - timeout=30, - ) - - logs = logs_result.stdout + logs_result.stderr - - # Check for network error patterns (case-insensitive) - network_errors = [] - for pattern in self.NETWORK_ERROR_PATTERNS: - matches = re.findall(pattern, logs, re.IGNORECASE) - if matches: - network_errors.extend(matches) - - # Check if container is still running - inspect_result = subprocess.run( - [ - "docker", - "inspect", - "-f", - "{{.State.Running}}", - self.TEST_CONTAINER_NAME, - ], - capture_output=True, - text=True, - timeout=10, - ) - - is_running = inspect_result.stdout.strip() == "true" - - # If container crashed, get exit code and reason - - # If container crashed, get exit code and reason - if not is_running: - exit_result = subprocess.run( - [ - "docker", - "inspect", - "-f", - "{{.State.ExitCode}}", - self.TEST_CONTAINER_NAME, - ], - capture_output=True, - text=True, - timeout=10, - ) - exit_code = exit_result.stdout.strip() - - # Container not running is OK if it didn't crash due to network issues - # Check if exit was due to network errors - if network_errors: - pytest.fail( - f"Container failed with network errors (exit code {exit_code}): " - f"{network_errors}\n\nFull logs:\n{logs}" - ) - - # Assert no network errors were found - assert len(network_errors) == 0, ( - f"Container made network requests during startup that failed: " - f"{network_errors}" - ) - - def test_no_external_urls_in_startup_code(self): - """ - Static analysis test: check that startup code doesn't contain - hardcoded external URLs that would be called during import/startup. - - This is a complementary test to catch issues without needing Docker. - """ - # Directories to check for startup code - startup_dirs = [ - "litellm/proxy", - "litellm/__init__.py", - "litellm/main.py", - ] - - # Patterns that indicate external URL calls during startup (not in functions) - problematic_patterns = [ - # Immediate HTTP calls (not inside functions) - r'^requests\.get\(["\'](https?://)', - r'^httpx\.get\(["\'](https?://)', - r'^urllib\.request\.urlopen\(["\'](https?://)', - ] - - # Files that are OK to have URLs (they're called on-demand, not startup) - allowed_files = [ - "model_prices_and_context_window.json", # Static data file - "test_", # Test files - "_test.py", - "conftest.py", - ] - - workspace_root = os.path.dirname(os.path.dirname(os.path.dirname(__file__))) - - issues_found = [] - - for startup_dir in startup_dirs: - full_path = os.path.join(workspace_root, startup_dir) - if not os.path.exists(full_path): - continue - - if os.path.isfile(full_path): - files_to_check = [full_path] - else: - files_to_check = [] - for root, _dirs, files in os.walk(full_path): - for f in files: - if f.endswith(".py"): - files_to_check.append(os.path.join(root, f)) - - for filepath in files_to_check: - # Skip allowed files - if any(allowed in filepath for allowed in allowed_files): - continue - - try: - with open(filepath, "r") as f: - content = f.read() - - for pattern in problematic_patterns: - matches = re.findall(pattern, content, re.MULTILINE) - if matches: - issues_found.append(f"{filepath}: {pattern} matched") - except Exception: - pass # Skip unreadable files - - # This test is informational - we document but don't fail - if issues_found: - pass - - -@pytest.mark.skipif( - not is_docker_available(), - reason="Docker not available", -) -def test_container_build_no_network_fetch(): - """ - Test that the Docker build process doesn't require network for runtime. - - This verifies that all dependencies are properly bundled and no - runtime network calls are made during container initialization. - - Note: Build itself may need network to resolve dependencies, but runtime should not. - """ - # This is a simplified version - full test would need to: - # 1. Build image with --network=none (requires pre-cached deps) - # 2. Or run built image in isolated network - - # For now, just verify the Dockerfile doesn't have wget/curl in CMD/ENTRYPOINT - workspace_root = os.path.dirname(os.path.dirname(os.path.dirname(__file__))) - dockerfile_path = os.path.join(workspace_root, "Dockerfile") - - if not os.path.exists(dockerfile_path): - pytest.skip("Dockerfile not found") - - with open(dockerfile_path, "r") as f: - content = f.read() - - # Check for network calls in CMD/ENTRYPOINT - problematic = [] - lines = content.split("\n") - for i, line in enumerate(lines, 1): - line_upper = line.strip().upper() - if line_upper.startswith(("CMD", "ENTRYPOINT")): - if any( - cmd in line.lower() - for cmd in ["curl", "wget", "fetch", "http://", "https://"] - ): - problematic.append(f"Line {i}: {line.strip()}") - - assert ( - len(problematic) == 0 - ), f"Dockerfile CMD/ENTRYPOINT contains network calls: {problematic}" +""" +Test to ensure Docker container does not go out to network on deploy. + +This test verifies that the LiteLLM proxy container does not make outbound +network requests during startup. This is important for: +1. Air-gapped environments where outbound network is restricted +2. Security compliance requiring no unexpected network calls +3. Fast container startup without network dependencies + +The test works by: +1. Building/running the container with network disabled +2. Verifying the container starts successfully without network +3. Checking for any errors related to network failures during startup +""" + +import os +import re +import subprocess +import time + +import pytest + +from tests._master_key import MASTER_KEY + + +def is_docker_available() -> bool: + """Check if Docker is available and running.""" + try: + result = subprocess.run( + ["docker", "info"], + capture_output=True, + text=True, + timeout=10, + ) + return result.returncode == 0 + except (subprocess.TimeoutExpired, FileNotFoundError): + return False + + +@pytest.mark.skipif( + not is_docker_available(), + reason="Docker not available", +) +class TestDockerNoNetworkOnDeploy: + """ + Test suite for verifying Docker container starts without network access. + """ + + # Container and image names for testing + TEST_IMAGE_NAME = "litellm-no-network-test" + TEST_CONTAINER_NAME = "litellm-no-network-test-container" + + # Timeout for container operations + CONTAINER_START_TIMEOUT = 60 # seconds + + # Patterns that indicate network-related failures during startup + NETWORK_ERROR_PATTERNS = [ + r"connection refused", + r"network is unreachable", + r"could not resolve host", + r"name resolution failed", + r"dns lookup failed", + r"failed to establish.*connection", + r"socket.gaierror", + r"urllib.error.URLError.*Errno", + r"requests.exceptions.ConnectionError", + r"httpx.*ConnectError", + r"aiohttp.*ClientConnectorError", + ] + + # Patterns that indicate EXPECTED local-only startup + LOCAL_STARTUP_PATTERNS = [ + r"starting.*(server|proxy)", + r"listening on", + r"uvicorn running", + r"application startup complete", + ] + + @pytest.fixture(autouse=True) + def cleanup(self): + """Cleanup any existing test containers before and after each test.""" + self._cleanup_container() + yield + self._cleanup_container() + + def _cleanup_container(self): + """Remove test container if it exists.""" + subprocess.run( + ["docker", "rm", "-f", self.TEST_CONTAINER_NAME], + capture_output=True, + timeout=30, + ) + + def _build_test_image(self) -> bool: + """ + Build the Docker image if needed. + Returns True if image is available (built or already exists). + """ + # Check if main litellm image exists + result = subprocess.run( + ["docker", "images", "-q", "litellm/litellm"], + capture_output=True, + text=True, + timeout=10, + ) + if result.stdout.strip(): + # Image exists, use it + return True + + # Try to build from Dockerfile + dockerfile_path = os.path.join( + os.path.dirname(__file__), + "..", + "..", + "..", + "Dockerfile", + ) + if os.path.exists(dockerfile_path): + result = subprocess.run( + [ + "docker", + "build", + "-t", + self.TEST_IMAGE_NAME, + "-f", + dockerfile_path, + os.path.dirname(dockerfile_path), + ], + capture_output=True, + text=True, + timeout=600, # 10 minutes for build + ) + return result.returncode == 0 + + return False + + def test_container_starts_without_network(self): + """ + Test that the container can start with network completely disabled. + + This test runs the container with --network=none to ensure no outbound + network requests are required during startup. + """ + # Use a minimal config that doesn't require external services + minimal_config = f""" +model_list: + - model_name: fake-model + litellm_params: + model: fake/fake-model + +general_settings: + master_key: {MASTER_KEY} + database_url: null + +environment_variables: {{}} +""" + + # Create a temporary config file + config_path = "/tmp/litellm_test_config.yaml" + with open(config_path, "w") as f: + f.write(minimal_config) + + image_to_use = "litellm/litellm" + # Check if image exists + result = subprocess.run( + ["docker", "images", "-q", image_to_use], + capture_output=True, + text=True, + timeout=10, + ) + if not result.stdout.strip(): + pytest.skip(f"Docker image {image_to_use} not available") + + # Run container with network disabled + run_cmd = [ + "docker", + "run", + "--name", + self.TEST_CONTAINER_NAME, + "--network=none", # Disable all network access + "-v", + f"{config_path}:/app/config.yaml:ro", + "-e", + f"LITELLM_MASTER_KEY={MASTER_KEY}", + "-e", + "DATABASE_URL=", # Empty to disable DB + "-e", + "STORE_MODEL_IN_DB=false", + "-e", + "LITELLM_LOG=DEBUG", + "-d", # Detached mode + image_to_use, + "--config", + "/app/config.yaml", + ] + + result = subprocess.run( + run_cmd, + capture_output=True, + text=True, + timeout=30, + ) + + if result.returncode != 0: + pytest.fail(f"Failed to start container: {result.stderr}") + + # Wait for container to start up + time.sleep(5) + + # Check container logs for any network-related errors + logs_result = subprocess.run( + ["docker", "logs", self.TEST_CONTAINER_NAME], + capture_output=True, + text=True, + timeout=30, + ) + + logs = logs_result.stdout + logs_result.stderr + + # Check for network error patterns (case-insensitive) + network_errors = [] + for pattern in self.NETWORK_ERROR_PATTERNS: + matches = re.findall(pattern, logs, re.IGNORECASE) + if matches: + network_errors.extend(matches) + + # Check if container is still running + inspect_result = subprocess.run( + [ + "docker", + "inspect", + "-f", + "{{.State.Running}}", + self.TEST_CONTAINER_NAME, + ], + capture_output=True, + text=True, + timeout=10, + ) + + is_running = inspect_result.stdout.strip() == "true" + + # If container crashed, get exit code and reason + + # If container crashed, get exit code and reason + if not is_running: + exit_result = subprocess.run( + [ + "docker", + "inspect", + "-f", + "{{.State.ExitCode}}", + self.TEST_CONTAINER_NAME, + ], + capture_output=True, + text=True, + timeout=10, + ) + exit_code = exit_result.stdout.strip() + + # Container not running is OK if it didn't crash due to network issues + # Check if exit was due to network errors + if network_errors: + pytest.fail( + f"Container failed with network errors (exit code {exit_code}): " + f"{network_errors}\n\nFull logs:\n{logs}" + ) + + # Assert no network errors were found + assert len(network_errors) == 0, ( + f"Container made network requests during startup that failed: " + f"{network_errors}" + ) + + def test_no_external_urls_in_startup_code(self): + """ + Static analysis test: check that startup code doesn't contain + hardcoded external URLs that would be called during import/startup. + + This is a complementary test to catch issues without needing Docker. + """ + # Directories to check for startup code + startup_dirs = [ + "litellm/proxy", + "litellm/__init__.py", + "litellm/main.py", + ] + + # Patterns that indicate external URL calls during startup (not in functions) + problematic_patterns = [ + # Immediate HTTP calls (not inside functions) + r'^requests\.get\(["\'](https?://)', + r'^httpx\.get\(["\'](https?://)', + r'^urllib\.request\.urlopen\(["\'](https?://)', + ] + + # Files that are OK to have URLs (they're called on-demand, not startup) + allowed_files = [ + "model_prices_and_context_window.json", # Static data file + "test_", # Test files + "_test.py", + "conftest.py", + ] + + workspace_root = os.path.dirname(os.path.dirname(os.path.dirname(__file__))) + + issues_found = [] + + for startup_dir in startup_dirs: + full_path = os.path.join(workspace_root, startup_dir) + if not os.path.exists(full_path): + continue + + if os.path.isfile(full_path): + files_to_check = [full_path] + else: + files_to_check = [] + for root, _dirs, files in os.walk(full_path): + for f in files: + if f.endswith(".py"): + files_to_check.append(os.path.join(root, f)) + + for filepath in files_to_check: + # Skip allowed files + if any(allowed in filepath for allowed in allowed_files): + continue + + try: + with open(filepath, "r") as f: + content = f.read() + + for pattern in problematic_patterns: + matches = re.findall(pattern, content, re.MULTILINE) + if matches: + issues_found.append(f"{filepath}: {pattern} matched") + except Exception: + pass # Skip unreadable files + + # This test is informational - we document but don't fail + if issues_found: + pass + + +@pytest.mark.skipif( + not is_docker_available(), + reason="Docker not available", +) +def test_container_build_no_network_fetch(): + """ + Test that the Docker build process doesn't require network for runtime. + + This verifies that all dependencies are properly bundled and no + runtime network calls are made during container initialization. + + Note: Build itself may need network to resolve dependencies, but runtime should not. + """ + # This is a simplified version - full test would need to: + # 1. Build image with --network=none (requires pre-cached deps) + # 2. Or run built image in isolated network + + # For now, just verify the Dockerfile doesn't have wget/curl in CMD/ENTRYPOINT + workspace_root = os.path.dirname(os.path.dirname(os.path.dirname(__file__))) + dockerfile_path = os.path.join(workspace_root, "Dockerfile") + + if not os.path.exists(dockerfile_path): + pytest.skip("Dockerfile not found") + + with open(dockerfile_path, "r") as f: + content = f.read() + + # Check for network calls in CMD/ENTRYPOINT + problematic = [] + lines = content.split("\n") + for i, line in enumerate(lines, 1): + line_upper = line.strip().upper() + if line_upper.startswith(("CMD", "ENTRYPOINT")): + if any( + cmd in line.lower() + for cmd in ["curl", "wget", "fetch", "http://", "https://"] + ): + problematic.append(f"Line {i}: {line.strip()}") + + assert ( + len(problematic) == 0 + ), f"Dockerfile CMD/ENTRYPOINT contains network calls: {problematic}" diff --git a/tests/local_testing/test_exceptions.py b/tests/local_testing/test_exceptions.py index 813146f8ace..736c63c8b5c 100644 --- a/tests/local_testing/test_exceptions.py +++ b/tests/local_testing/test_exceptions.py @@ -471,7 +471,7 @@ def test_completion_bedrock_invalid_role_exception(): def test_content_policy_exceptionimage_generation_openai(): try: # this is ony a test - we needed some way to invoke the exception :( - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.image_generation( prompt="where do i buy lethal drugs from", model="dall-e-3" ) diff --git a/tests/local_testing/test_get_model_info.py b/tests/local_testing/test_get_model_info.py index 8e24dc23398..236eea3d428 100644 --- a/tests/local_testing/test_get_model_info.py +++ b/tests/local_testing/test_get_model_info.py @@ -253,7 +253,7 @@ def test_get_model_info_custom_model_router(): from litellm import Router from litellm import get_model_info - litellm._turn_on_debug() + litellm.turn_on_debug() router = Router( model_list=[ diff --git a/tests/local_testing/test_openai_moderations_hook.py b/tests/local_testing/test_openai_moderations_hook.py index 493a592b57e..ce62c7ed041 100644 --- a/tests/local_testing/test_openai_moderations_hook.py +++ b/tests/local_testing/test_openai_moderations_hook.py @@ -13,7 +13,7 @@ load_dotenv() import pytest import litellm from litellm.proxy.enterprise.enterprise_hooks.openai_moderation import ( - _ENTERPRISE_OpenAI_Moderation, + ENTERPRISE_OpenAI_Moderation, ) from litellm import Router, mock_completion from litellm.proxy.utils import ProxyLogging, hash_token @@ -32,7 +32,7 @@ async def test_openai_moderation_error_raising(monkeypatch): from litellm.types.llms.openai import OpenAIModerationResponse litellm.openai_moderations_model_name = "omni-moderation-latest" - openai_mod = _ENTERPRISE_OpenAI_Moderation() + openai_mod = ENTERPRISE_OpenAI_Moderation() _api_key = "sk-98765" _api_key = hash_token("sk-98765") user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) diff --git a/tests/local_testing/test_router.py b/tests/local_testing/test_router.py index 4965fa631a9..ea151d6f65b 100644 --- a/tests/local_testing/test_router.py +++ b/tests/local_testing/test_router.py @@ -24,8 +24,8 @@ from litellm import Router from litellm.router import Deployment, LiteLLM_Params from litellm.types.router import ModelInfo from litellm.router_utils.cooldown_handlers import ( - _async_get_cooldown_deployments, - _get_cooldown_deployments, + async_get_cooldown_deployments, + get_cooldown_deployments, ) from litellm.types.router import DeploymentTypedDict @@ -496,7 +496,7 @@ async def test_async_router_context_window_fallback(sync_mode): from large_text import text litellm.set_verbose = False - litellm._turn_on_debug() + litellm.turn_on_debug() print(f"len(text): {len(text)}") try: @@ -1881,11 +1881,11 @@ async def test_aaarouter_dynamic_cooldown_message_retry_time(sync_mode): ) if sync_mode: - cooldown_deployments = _get_cooldown_deployments( + cooldown_deployments = get_cooldown_deployments( litellm_router_instance=router, parent_otel_span=None ) else: - cooldown_deployments = await _async_get_cooldown_deployments( + cooldown_deployments = await async_get_cooldown_deployments( litellm_router_instance=router, parent_otel_span=None ) diff --git a/tests/local_testing/test_router_batch_completion.py b/tests/local_testing/test_router_batch_completion.py index 6fd89065c1d..67e6acd51b8 100644 --- a/tests/local_testing/test_router_batch_completion.py +++ b/tests/local_testing/test_router_batch_completion.py @@ -124,7 +124,7 @@ async def test_batch_completion_fastest_response_unit_test(): @pytest.mark.asyncio async def test_batch_completion_fastest_response_streaming(): litellm.set_verbose = True - litellm._turn_on_debug() + litellm.turn_on_debug() router = litellm.Router( model_list=[ diff --git a/tests/local_testing/test_router_cooldown_handlers.py b/tests/local_testing/test_router_cooldown_handlers.py index 0ec9623538a..72621414b95 100644 --- a/tests/local_testing/test_router_cooldown_handlers.py +++ b/tests/local_testing/test_router_cooldown_handlers.py @@ -19,7 +19,7 @@ import litellm from litellm import Router from litellm.integrations.custom_logger import CustomLogger from litellm.router_utils.cooldown_handlers import ( - _async_get_cooldown_deployments, + async_get_cooldown_deployments, _should_run_cooldown_logic, ) from litellm.types.router import ( @@ -163,7 +163,7 @@ async def test_cooldown_time_zero_uses_zero_not_default(): mock_add_cooldown.assert_not_called() # Also verify the deployment is not in cooldown - cooldown_list = await _async_get_cooldown_deployments( + cooldown_list = await async_get_cooldown_deployments( litellm_router_instance=router, parent_otel_span=None ) assert len(cooldown_list) == 0 @@ -474,7 +474,7 @@ async def test_single_deployment_no_cooldowns_test_prod_mock_completion_calls(): except litellm.RateLimitError: pass - cooldown_list = await _async_get_cooldown_deployments( + cooldown_list = await async_get_cooldown_deployments( litellm_router_instance=router, parent_otel_span=None ) assert len(cooldown_list) == 0 @@ -582,7 +582,7 @@ async def test_high_traffic_cooldowns_all_healthy_deployments(): raise e print("model_stats: ", model_stats) - cooldown_list = await _async_get_cooldown_deployments( + cooldown_list = await async_get_cooldown_deployments( litellm_router_instance=router, parent_otel_span=None ) assert len(cooldown_list) == 0 @@ -679,7 +679,7 @@ async def test_high_traffic_cooldowns_one_bad_deployment(): raise e print("model_stats: ", model_stats) - cooldown_list = await _async_get_cooldown_deployments( + cooldown_list = await async_get_cooldown_deployments( litellm_router_instance=router, parent_otel_span=None ) assert len(cooldown_list) == 1 @@ -779,7 +779,7 @@ async def test_high_traffic_cooldowns_one_rate_limited_deployment(): raise e print("model_stats: ", model_stats) - cooldown_list = await _async_get_cooldown_deployments( + cooldown_list = await async_get_cooldown_deployments( litellm_router_instance=router, parent_otel_span=None ) assert len(cooldown_list) == 1 @@ -788,7 +788,7 @@ async def test_high_traffic_cooldowns_one_rate_limited_deployment(): """ Unit tests for router set_cooldowns -1. _set_cooldown_deployments() will cooldown a deployment after it fails 50% requests +1. set_cooldown_deployments() will cooldown a deployment after it fails 50% requests """ @@ -836,7 +836,7 @@ async def test_router_fallbacks_with_cooldowns_and_dynamic_credentials(): A 429 answered to a caller-supplied credential cools down none of the shared deployments, so the next credential still reaches them, while a 429 owned by a shared deployment does """ - from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments + from litellm.router_utils.cooldown_handlers import async_get_cooldown_deployments router = Router( model_list=[ @@ -856,7 +856,7 @@ async def test_router_fallbacks_with_cooldowns_and_dynamic_credentials(): model="gpt-3.5-turbo", messages=messages, api_key="my-bad-key-1", mock_response="litellm.RateLimitError" ) await asyncio.sleep(1) - assert await _async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) == [] + assert await async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) == [] response = await router.acompletion( model="gpt-3.5-turbo", messages=messages, api_key="my-good-key-2", mock_response="served with credential 2" @@ -866,5 +866,5 @@ async def test_router_fallbacks_with_cooldowns_and_dynamic_credentials(): with pytest.raises(litellm.RateLimitError): await router.acompletion(model="gpt-3.5-turbo", messages=messages, mock_response="litellm.RateLimitError") await asyncio.sleep(1) - cooled_down = await _async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) + cooled_down = await async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) assert len(cooled_down) == 1 and cooled_down[0] in {"123", "456"} diff --git a/tests/local_testing/test_router_pattern_matching.py b/tests/local_testing/test_router_pattern_matching.py index 6ffc5316f2e..82a7f851a73 100644 --- a/tests/local_testing/test_router_pattern_matching.py +++ b/tests/local_testing/test_router_pattern_matching.py @@ -99,9 +99,9 @@ def test_pattern_to_regex(): Tests that the pattern is converted to a regex """ router = PatternMatchRouter() - assert router._pattern_to_regex("openai/*") == "openai/(.*)" + assert router.pattern_to_regex("openai/*") == "openai/(.*)" assert ( - router._pattern_to_regex("openai/fo::*::static::*") + router.pattern_to_regex("openai/fo::*::static::*") == "openai/fo::(.*)::static::(.*)" ) diff --git a/tests/local_testing/test_router_utils.py b/tests/local_testing/test_router_utils.py index aa617b09731..8d30da1a180 100644 --- a/tests/local_testing/test_router_utils.py +++ b/tests/local_testing/test_router_utils.py @@ -317,7 +317,7 @@ def test_router_get_model_access_groups(potential_access_group, expected_result) }, ] ) - access_groups = router._is_model_access_group_for_wildcard_route( + access_groups = router.is_model_access_group_for_wildcard_route( model_access_group=potential_access_group ) assert access_groups == expected_result @@ -330,7 +330,7 @@ def test_router_redis_cache(): redis_cache = MagicMock() - router._update_redis_cache(cache=redis_cache) + router.update_redis_cache(cache=redis_cache) assert router.cache.redis_cache == redis_cache diff --git a/tests/local_testing/test_secret_detect_hook.py b/tests/local_testing/test_secret_detect_hook.py index 04278137d4e..137d6c16101 100644 --- a/tests/local_testing/test_secret_detect_hook.py +++ b/tests/local_testing/test_secret_detect_hook.py @@ -262,7 +262,7 @@ async def test_chat_completion_request_with_redaction(): setattr(proxy_server, "llm_router", router) _test_logger = testLogger() litellm.callbacks = [_ENTERPRISE_SecretDetection(), _test_logger] - litellm._turn_on_debug() + litellm.turn_on_debug() # Prepare the query string query_params = "param1=value1¶m2=value2" diff --git a/tests/local_testing/test_streaming.py b/tests/local_testing/test_streaming.py index 6e102b89554..ad1b35a4c18 100644 --- a/tests/local_testing/test_streaming.py +++ b/tests/local_testing/test_streaming.py @@ -508,7 +508,7 @@ def test_completion_model_stream(model): @pytest.mark.flaky(retries=3, delay=1) async def test_completion_gemini_stream(sync_mode): try: - litellm._turn_on_debug() + litellm.turn_on_debug() print("Streaming gemini response") function1 = [ { @@ -3338,7 +3338,7 @@ def test_mock_response_iterator_tool_use(): def test_reasoning_content_completion(model): # litellm.set_verbose = True try: - # litellm._turn_on_debug() + # litellm.turn_on_debug() resp = litellm.completion( model=model, messages=[{"role": "user", "content": "Tell me a joke."}], diff --git a/tests/local_testing/test_text_completion.py b/tests/local_testing/test_text_completion.py index ea34b2dd21a..eaf80374687 100644 --- a/tests/local_testing/test_text_completion.py +++ b/tests/local_testing/test_text_completion.py @@ -4146,7 +4146,7 @@ def test_completion_vllm(provider): @pytest.mark.skip(reason="fireworks is having an active outage") def test_completion_fireworks_ai_multiple_choices(): - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.text_completion( model="fireworks_ai/llama-v3p1-8b-instruct", prompt=["halo", "hi", "halo", "hi"], diff --git a/tests/logging_callback_tests/test_alerting.py b/tests/logging_callback_tests/test_alerting.py index de84443814c..38539af76f9 100644 --- a/tests/logging_callback_tests/test_alerting.py +++ b/tests/logging_callback_tests/test_alerting.py @@ -739,7 +739,7 @@ async def test_langfuse_trace_id(): - Unit test for `_add_langfuse_trace_id_to_alert` function in slack_alerting.py """ from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert + from litellm.integrations.SlackAlerting.utils import add_langfuse_trace_id_to_alert litellm.success_callback = ["langfuse"] @@ -762,7 +762,7 @@ async def test_langfuse_trace_id(): await asyncio.sleep(3) - assert litellm_logging_obj._get_trace_id(service_name="langfuse") is not None + assert litellm_logging_obj.get_trace_id(service_name="langfuse") is not None slack_alerting = SlackAlerting( alerting_threshold=32, @@ -771,7 +771,7 @@ async def test_langfuse_trace_id(): internal_usage_cache=DualCache(), ) - trace_url = await _add_langfuse_trace_id_to_alert( + trace_url = await add_langfuse_trace_id_to_alert( request_data={"litellm_logging_obj": litellm_logging_obj} ) @@ -779,7 +779,7 @@ async def test_langfuse_trace_id(): returned_trace_id = trace_url.split("/")[-1] - assert returned_trace_id == litellm_logging_obj._get_trace_id( + assert returned_trace_id == litellm_logging_obj.get_trace_id( service_name="langfuse" ) diff --git a/tests/logging_callback_tests/test_assemble_streaming_responses.py b/tests/logging_callback_tests/test_assemble_streaming_responses.py index d6905ce3565..ee3397b3567 100644 --- a/tests/logging_callback_tests/test_assemble_streaming_responses.py +++ b/tests/logging_callback_tests/test_assemble_streaming_responses.py @@ -29,7 +29,7 @@ from litellm import ( ) from litellm.litellm_core_utils.logging_utils import ( - _assemble_complete_response_from_streaming_chunks, + assemble_complete_response_from_streaming_chunks, ) @@ -66,7 +66,7 @@ def test_assemble_complete_response_from_streaming_chunks_1(is_async): "usage": None, } chunk = ModelResponseStream(**chunk) - complete_streaming_response = _assemble_complete_response_from_streaming_chunks( + complete_streaming_response = assemble_complete_response_from_streaming_chunks( result=chunk, start_time=datetime.now(), end_time=datetime.now(), @@ -106,7 +106,7 @@ def test_assemble_complete_response_from_streaming_chunks_1(is_async): "usage": None, } chunk = ModelResponseStream(**chunk) - complete_streaming_response = _assemble_complete_response_from_streaming_chunks( + complete_streaming_response = assemble_complete_response_from_streaming_chunks( result=chunk, start_time=datetime.now(), end_time=datetime.now(), @@ -169,7 +169,7 @@ def test_assemble_complete_response_from_streaming_chunks_2(is_async): chunk = ModelResponseStream(**chunk) chunk = _text_completion_stream_wrapper.convert_to_text_completion_object(chunk) - complete_streaming_response = _assemble_complete_response_from_streaming_chunks( + complete_streaming_response = assemble_complete_response_from_streaming_chunks( result=chunk, start_time=datetime.now(), end_time=datetime.now(), @@ -210,7 +210,7 @@ def test_assemble_complete_response_from_streaming_chunks_2(is_async): } chunk = ModelResponseStream(**chunk) chunk = _text_completion_stream_wrapper.convert_to_text_completion_object(chunk) - complete_streaming_response = _assemble_complete_response_from_streaming_chunks( + complete_streaming_response = assemble_complete_response_from_streaming_chunks( result=chunk, start_time=datetime.now(), end_time=datetime.now(), @@ -264,7 +264,7 @@ def test_assemble_complete_response_from_streaming_chunks_3(is_async): "usage": None, } chunk = ModelResponseStream(**chunk) - complete_streaming_response = _assemble_complete_response_from_streaming_chunks( + complete_streaming_response = assemble_complete_response_from_streaming_chunks( result=chunk, start_time=datetime.now(), end_time=datetime.now(), @@ -284,7 +284,7 @@ def test_assemble_complete_response_from_streaming_chunks_3(is_async): # now add a chunk to the 2nd list - complete_streaming_response = _assemble_complete_response_from_streaming_chunks( + complete_streaming_response = assemble_complete_response_from_streaming_chunks( result=chunk, start_time=datetime.now(), end_time=datetime.now(), @@ -345,7 +345,7 @@ def test_assemble_complete_response_from_streaming_chunks_4(is_async): # remove attribute id from chunk del chunk.object - complete_streaming_response = _assemble_complete_response_from_streaming_chunks( + complete_streaming_response = assemble_complete_response_from_streaming_chunks( result=chunk, start_time=datetime.now(), end_time=datetime.now(), diff --git a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py index 5571c15beff..c8c87c21010 100644 --- a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py +++ b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py @@ -117,7 +117,7 @@ async def test_vector_store_hook_routes_search_through_proxy_router( async def test_e2e_bedrock_knowledgebase_retrieval_with_completion( setup_vector_store_registry, ): - litellm._turn_on_debug() + litellm.turn_on_debug() client = AsyncHTTPHandler() print("value of litellm.vector_store_registry:", litellm.vector_store_registry) @@ -187,7 +187,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call( """ # Init client - litellm._turn_on_debug() + litellm.turn_on_debug() async_client = AsyncHTTPHandler() response = await litellm.acompletion( model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", @@ -231,7 +231,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_streaming( """ # Init client - # litellm._turn_on_debug() + # litellm.turn_on_debug() async_client = AsyncHTTPHandler() response = await litellm.acompletion( model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}", @@ -290,7 +290,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools( """ # Init client - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.acompletion( model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}", messages=[{"role": "user", "content": "what is litellm?"}], @@ -310,7 +310,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools_ In this case we filter for a non-existent user_id, which should return no results. """ - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.acompletion( model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}", @@ -731,7 +731,7 @@ async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_regist async def test_e2e_bedrock_knowledgebase_retrieval_without_vector_store_registry( setup_vector_store_registry, ): - litellm._turn_on_debug() + litellm.turn_on_debug() client = AsyncHTTPHandler() litellm.vector_store_registry = None @@ -796,7 +796,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_vector_store_not_in_regi In this test newUnknownVectorStoreId is not in the registry, so no vector store request is made """ - litellm._turn_on_debug() + litellm.turn_on_debug() client = AsyncHTTPHandler() if litellm.vector_store_registry is not None: diff --git a/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py b/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py index 53fe493ad9f..2d3933c3dad 100644 --- a/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py +++ b/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py @@ -45,7 +45,7 @@ class TestCustomLogger(CustomLogger): async def _setup_web_search_test(): """Helper function to setup common test requirements""" - litellm._turn_on_debug() + litellm.turn_on_debug() test_custom_logger = TestCustomLogger() litellm.callbacks = [test_custom_logger] return test_custom_logger diff --git a/tests/logging_callback_tests/test_custom_callback_router.py b/tests/logging_callback_tests/test_custom_callback_router.py index 8cbe5fc6ccc..97def5e81b0 100644 --- a/tests/logging_callback_tests/test_custom_callback_router.py +++ b/tests/logging_callback_tests/test_custom_callback_router.py @@ -281,7 +281,7 @@ class CompletionCustomHandler( assert isinstance(kwargs["model"], str) # checking we use base_model for azure cost calculation - base_model = litellm.utils._get_base_model_from_metadata( + base_model = litellm.utils.get_base_model_from_metadata( model_call_details=kwargs ) diff --git a/tests/logging_callback_tests/test_langfuse_e2e_test.py b/tests/logging_callback_tests/test_langfuse_e2e_test.py index 61a93a175d9..8330a18490d 100644 --- a/tests/logging_callback_tests/test_langfuse_e2e_test.py +++ b/tests/logging_callback_tests/test_langfuse_e2e_test.py @@ -367,7 +367,7 @@ class TestLangfuseLogging: async def test_langfuse_logging_completion_with_malformed_llm_response(self, mock_setup): """Test Langfuse logging for chat completion with malformed LLM response""" setup = mock_setup - litellm._turn_on_debug() + litellm.turn_on_debug() with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): mock_response = litellm.ModelResponse( choices=[], @@ -393,7 +393,7 @@ class TestLangfuseLogging: async def test_langfuse_logging_completion_with_bedrock_llm_response(self, mock_setup): """Test Langfuse logging for chat completion with malformed LLM response""" setup = mock_setup - litellm._turn_on_debug() + litellm.turn_on_debug() with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): mock_response = litellm.ModelResponse( choices=[], @@ -426,7 +426,7 @@ class TestLangfuseLogging: async def test_langfuse_logging_completion_with_vertex_llm_response(self, mock_setup): """Test Langfuse logging for chat completion with malformed LLM response""" setup = mock_setup - litellm._turn_on_debug() + litellm.turn_on_debug() with patch(LANGFUSE_EXPORT_POST, setup["mock_post"]): mock_response = litellm.ModelResponse( choices=[], @@ -507,7 +507,7 @@ class TestLangfuseLogging: @pytest.mark.flaky(retries=3, delay=1) async def test_langfuse_logging_with_router(self, mock_setup): """Test Langfuse logging with router""" - litellm._turn_on_debug() + litellm.turn_on_debug() router = litellm.Router( model_list=[ { diff --git a/tests/logging_callback_tests/test_langfuse_unit_tests.py b/tests/logging_callback_tests/test_langfuse_unit_tests.py index 61316204fc3..b7554cb129b 100644 --- a/tests/logging_callback_tests/test_langfuse_unit_tests.py +++ b/tests/logging_callback_tests/test_langfuse_unit_tests.py @@ -352,7 +352,7 @@ def test_langfuse_e2e_sync(monkeypatch): stream=False, mock_response="Hello from litellm 2", ) - for logger in litellm.logging_callback_manager._get_all_callbacks(): + for logger in litellm.logging_callback_manager.get_all_callbacks(): if isinstance(logger, LangFuseLogger): logger.flush() deadline = time.time() + 10 diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp_providers.py b/tests/mcp_tests/test_aresponses_api_with_mcp_providers.py index 72a0415cf88..1f9e3e4eb76 100644 --- a/tests/mcp_tests/test_aresponses_api_with_mcp_providers.py +++ b/tests/mcp_tests/test_aresponses_api_with_mcp_providers.py @@ -84,7 +84,7 @@ async def test_streaming_responses_api_with_mcp_tools( ) as mock_get_tools, patch.object( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", new_callable=AsyncMock, ) as mock_execute_tools, ): @@ -211,7 +211,7 @@ async def test_streaming_mcp_event_order_and_response_id_consistency( ) as mock_get_tools, patch.object( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", new_callable=AsyncMock, ) as mock_execute_tools, ): diff --git a/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py b/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py index 520b31513f5..a7f04466d14 100644 --- a/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py @@ -129,7 +129,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): correctly and the provider is creating a cache. """ _skip_live_prompt_caching_test() - litellm._turn_on_debug() + litellm.turn_on_debug() messages = self.get_messages_with_cache_control() @@ -167,7 +167,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): This validates that caching is working end-to-end. """ _skip_live_prompt_caching_test() - litellm._turn_on_debug() + litellm.turn_on_debug() messages = self.get_messages_with_cache_control() @@ -208,7 +208,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): E2E test: Prompt caching with system message should work. """ _skip_live_prompt_caching_test() - litellm._turn_on_debug() + litellm.turn_on_debug() messages = [ { @@ -270,7 +270,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): are correctly returned in the streaming response's message_delta event. """ _skip_live_prompt_caching_test() - litellm._turn_on_debug() + litellm.turn_on_debug() messages = self.get_messages_with_cache_control() @@ -368,7 +368,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): E2E test: Second streaming call should return cache_read_input_tokens > 0. """ _skip_live_prompt_caching_test() - litellm._turn_on_debug() + litellm.turn_on_debug() messages = self.get_messages_with_cache_control() @@ -447,7 +447,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): wasn't supported. """ _skip_live_prompt_caching_test() - litellm._turn_on_debug() + litellm.turn_on_debug() messages = self.get_messages_with_cache_control() diff --git a/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py b/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py index 6a5bf627ac7..9e706f99316 100644 --- a/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py @@ -124,7 +124,7 @@ class BaseAnthropicMessagesToolSearchTest(ABC): This validates that the tool search beta header is being passed via extra_headers and forwarded correctly to the downstream provider. """ - litellm._turn_on_debug() + litellm.turn_on_debug() tools = self.get_tools_with_tool_search() messages = [{"role": "user", "content": "What's the weather in San Francisco?"}] @@ -155,7 +155,7 @@ class BaseAnthropicMessagesToolSearchTest(ABC): This validates that when the user asks about weather, the model discovers the get_weather tool via tool search and attempts to use it. """ - litellm._turn_on_debug() + litellm.turn_on_debug() tools = self.get_tools_with_tool_search() messages = [ @@ -195,7 +195,7 @@ class BaseAnthropicMessagesToolSearchTest(ABC): """ E2E test: Tool search should work with streaming responses. """ - litellm._turn_on_debug() + litellm.turn_on_debug() tools = self.get_tools_with_tool_search() messages = [{"role": "user", "content": "What's the weather like in Tokyo?"}] @@ -243,7 +243,7 @@ class BaseAnthropicMessagesToolSearchTest(ABC): This validates that the model can discover the appropriate tool from a larger catalog of deferred tools. """ - litellm._turn_on_debug() + litellm.turn_on_debug() tools = self.get_tools_with_tool_search() messages = [ diff --git a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py index ae8404cd6f9..147762a19a4 100644 --- a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py @@ -67,7 +67,7 @@ class BaseAnthropicMessagesTest: @pytest.mark.asyncio async def test_non_streaming_base(self): """Base test for non-streaming requests""" - litellm._turn_on_debug() + litellm.turn_on_debug() request_params = self.model_config @@ -134,7 +134,7 @@ class BaseAnthropicMessagesTest: Issue: https://github.com/BerriAI/litellm/issues/20342 """ - litellm._turn_on_debug() + litellm.turn_on_debug() request_params = self.model_config @@ -196,7 +196,7 @@ class BaseAnthropicMessagesTest: """ test_custom_logger = TestCustomLogger() litellm.callbacks = [test_custom_logger] - litellm._turn_on_debug() + litellm.turn_on_debug() router = Router( model_list=[ { diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py index ffbbf261e89..87a6d28f02f 100644 --- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py +++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py @@ -138,7 +138,7 @@ async def test_anthropic_messages_litellm_router_non_streaming(): """ Test the anthropic_messages with non-streaming request """ - litellm._turn_on_debug() + litellm.turn_on_debug() router = Router( model_list=[ { @@ -176,7 +176,7 @@ async def test_anthropic_messages_litellm_router_routing_strategy(): """ Test the anthropic_messages with routing strategy + non-streaming request """ - litellm._turn_on_debug() + litellm.turn_on_debug() router = Router( model_list=[ { @@ -218,7 +218,7 @@ async def test_anthropic_messages_fallbacks(): """ E2E test the anthropic_messages fallbacks from Anthropic API to Bedrock """ - litellm._turn_on_debug() + litellm.turn_on_debug() router = Router( model_list=[ { @@ -391,7 +391,7 @@ async def test_anthropic_messages_litellm_router_non_streaming_with_logging(): """ test_custom_logger = TestCustomLogger() litellm.callbacks = [test_custom_logger] - litellm._turn_on_debug() + litellm.turn_on_debug() MODEL_GROUP = "claude-special-alias" router = Router( model_list=[ @@ -807,7 +807,7 @@ def test_sync_openai_messages(): """ Test the anthropic_messages with sync request """ - litellm._turn_on_debug() + litellm.turn_on_debug() response = litellm.anthropic.messages.create( messages=[{"role": "user", "content": "Hello, can you tell me a short joke?"}], model="openai/gpt-4.1-mini", diff --git a/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py b/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py index e86c32f916d..451a09d30bb 100644 --- a/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py +++ b/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py @@ -20,7 +20,7 @@ async def test_anthropic_messages_litellm_router_bedrock(): Test the anthropic_messages with non-streaming request """ - litellm._turn_on_debug() + litellm.turn_on_debug() router = Router( model_list=[ { diff --git a/tests/pass_through_unit_tests/test_bedrock_tool_use_beta_header.py b/tests/pass_through_unit_tests/test_bedrock_tool_use_beta_header.py index a28b8a147af..675b8e6b5fb 100644 --- a/tests/pass_through_unit_tests/test_bedrock_tool_use_beta_header.py +++ b/tests/pass_through_unit_tests/test_bedrock_tool_use_beta_header.py @@ -19,7 +19,7 @@ async def test_bedrock_sonnet_4_5_with_advanced_tool_use_beta_header(): This should work without throwing "invalid beta flag" error because LiteLLM filters out the advanced-tool-use beta header for Bedrock Invoke API. """ - litellm._turn_on_debug() + litellm.turn_on_debug() response = await litellm.anthropic.messages.acreate( model="bedrock/invoke/us.anthropic.claude-sonnet-4-5-20250929-v1:0", messages=[{"role": "user", "content": "What is 2+2?"}], diff --git a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py b/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py index 8b3dc436b8f..a468c594aac 100644 --- a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py +++ b/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py @@ -48,7 +48,7 @@ class TestVertexAILivePassthroughLoggingHandler: """Create a mock logging object""" mock = MagicMock(spec=LiteLLMLoggingObj) mock.model_call_details = {} - mock._response_cost_calculator.return_value = None + mock.response_cost_calculator.return_value = None return mock @pytest.fixture @@ -788,7 +788,7 @@ class TestVertexAILivePassthroughIntegration: """Create a mock logging object""" mock = MagicMock(spec=LiteLLMLoggingObj) mock.model_call_details = {} - mock._response_cost_calculator.return_value = None + mock.response_cost_calculator.return_value = None return mock @patch( @@ -922,7 +922,7 @@ class TestVertexAILivePassthroughErrorHandling: """Create a mock logging object""" mock = MagicMock(spec=LiteLLMLoggingObj) mock.model_call_details = {} - mock._response_cost_calculator.return_value = None + mock.response_cost_calculator.return_value = None return mock def test_invalid_websocket_messages_format(self): diff --git a/tests/pass_through_unit_tests/test_websearch_interception_e2e.py b/tests/pass_through_unit_tests/test_websearch_interception_e2e.py index ca8f7baf01b..f9fa43a8f78 100644 --- a/tests/pass_through_unit_tests/test_websearch_interception_e2e.py +++ b/tests/pass_through_unit_tests/test_websearch_interception_e2e.py @@ -21,7 +21,7 @@ async def test_websearch_interception_non_streaming(): Test WebSearch interception with non-streaming request. Validates that agentic loop executes transparently. """ - litellm._turn_on_debug() + litellm.turn_on_debug() print("\n" + "=" * 80) print("E2E TEST 1: WebSearch Interception (Non-Streaming)") @@ -921,7 +921,7 @@ async def test_pre_request_hook_modifies_request_body(): from unittest.mock import AsyncMock, patch, MagicMock from litellm.constants import LITELLM_WEB_SEARCH_TOOL_NAME - litellm._turn_on_debug() + litellm.turn_on_debug() print("\n" + "=" * 80) print("UNIT TEST: Pre-Request Hook Modifies Request Body") diff --git a/tests/proxy_admin_ui_tests/test_sso_sign_in.py b/tests/proxy_admin_ui_tests/test_sso_sign_in.py index dd618cf3836..56778fc2171 100644 --- a/tests/proxy_admin_ui_tests/test_sso_sign_in.py +++ b/tests/proxy_admin_ui_tests/test_sso_sign_in.py @@ -58,7 +58,7 @@ async def test_auth_callback_new_user(mock_google_sso, mock_env_vars, prisma_cli from litellm._uuid import uuid import litellm - litellm._turn_on_debug() + litellm.turn_on_debug() # Generate a unique user ID unique_user_id = str(uuid.uuid4()) diff --git a/tests/router_unit_tests/test_router_batch_utils.py b/tests/router_unit_tests/test_router_batch_utils.py index 4336185a07f..6a73576ab42 100644 --- a/tests/router_unit_tests/test_router_batch_utils.py +++ b/tests/router_unit_tests/test_router_batch_utils.py @@ -6,7 +6,7 @@ from io import BytesIO from typing import Dict, List from litellm.router_utils.batch_utils import ( replace_model_in_jsonl, - _get_router_metadata_variable_name, + get_router_metadata_variable_name, InMemoryFile, parse_jsonl_with_embedded_newlines, ) @@ -100,16 +100,16 @@ def test_file_like_object(sample_file_like): def test_router_metadata_variable_name(): """Test that the variable name is correct""" - assert _get_router_metadata_variable_name(function_name="completion") == "metadata" + assert get_router_metadata_variable_name(function_name="completion") == "metadata" assert ( - _get_router_metadata_variable_name(function_name="batch") == "litellm_metadata" + get_router_metadata_variable_name(function_name="batch") == "litellm_metadata" ) assert ( - _get_router_metadata_variable_name(function_name="acreate_file") + get_router_metadata_variable_name(function_name="acreate_file") == "litellm_metadata" ) assert ( - _get_router_metadata_variable_name(function_name="aget_file") + get_router_metadata_variable_name(function_name="aget_file") == "litellm_metadata" ) diff --git a/tests/router_unit_tests/test_router_cooldown_per_deployment.py b/tests/router_unit_tests/test_router_cooldown_per_deployment.py index 964782c348b..5a30edc5c4f 100644 --- a/tests/router_unit_tests/test_router_cooldown_per_deployment.py +++ b/tests/router_unit_tests/test_router_cooldown_per_deployment.py @@ -329,7 +329,7 @@ class TestCooldownCacheTTLCorrection: class TestFallbackDeploymentCooldown: def test_trigger_cooldown_for_failed_deployment_calls_set_cooldown(self): """ - _trigger_cooldown_for_failed_deployment must call _set_cooldown_deployments + _trigger_cooldown_for_failed_deployment must call set_cooldown_deployments with the deployment ID stamped on the exception. """ mock_router = MagicMock() @@ -339,7 +339,7 @@ class TestFallbackDeploymentCooldown: exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") exc.failed_deployment_id = "fallback-deployment" - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment( litellm_router=mock_router, kwargs={}, @@ -354,11 +354,11 @@ class TestFallbackDeploymentCooldown: def test_trigger_cooldown_no_op_when_deployment_id_missing(self): """ _trigger_cooldown_for_failed_deployment must not raise and must skip - _set_cooldown_deployments when the exception has no failed_deployment_id. + set_cooldown_deployments when the exception has no failed_deployment_id. """ mock_router = MagicMock() - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment( litellm_router=mock_router, kwargs={}, @@ -386,7 +386,7 @@ class TestFallbackDeploymentCooldown: } } - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment( litellm_router=mock_router, kwargs=kwargs, @@ -409,7 +409,7 @@ class TestFallbackDeploymentCooldown: exc.failed_deployment_id = "fallback-deployment" with ( - patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown, + patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown, patch( "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" ) as mock_increment, @@ -433,7 +433,7 @@ class TestFallbackDeploymentCooldown: exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") exc.failed_deployment_id = "fallback-deployment" - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment( litellm_router=mock_router, kwargs={}, @@ -459,7 +459,7 @@ class TestFallbackDeploymentCooldown: exc.failed_deployment_id = "fallback-deployment" mark_advisor_orchestration_failure(exc) - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment( litellm_router=mock_router, kwargs={}, @@ -482,7 +482,7 @@ class TestFallbackDeploymentCooldown: exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") exc.failed_deployment_id = "fallback-deployment" - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment( litellm_router=mock_router, kwargs={}, @@ -505,7 +505,7 @@ class TestFallbackDeploymentCooldown: exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") exc.failed_deployment_id = "fallback-deployment" - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment( litellm_router=mock_router, kwargs={}, @@ -645,7 +645,7 @@ class TestDeploymentCallbackOnFailureCooldownTimePrecedence: }, } - with patch("litellm.router._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router.set_cooldown_deployments") as mock_set_cooldown: router.deployment_callback_on_failure( kwargs=kwargs, completion_response=None, @@ -681,7 +681,7 @@ class TestDeploymentCallbackOnFailureCooldownTimePrecedence: }, } - with patch("litellm.router._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router.set_cooldown_deployments") as mock_set_cooldown: router.deployment_callback_on_failure( kwargs=kwargs, completion_response=None, diff --git a/tests/router_unit_tests/test_router_endpoints.py b/tests/router_unit_tests/test_router_endpoints.py index a06aaa363ad..df2f352e6ea 100644 --- a/tests/router_unit_tests/test_router_endpoints.py +++ b/tests/router_unit_tests/test_router_endpoints.py @@ -1237,7 +1237,7 @@ async def test_init_containers_api_endpoints_managed_id_routes_via_generic_fallb ) router._ageneric_api_call_with_fallbacks = AsyncMock() - managed_id = ResponsesAPIRequestUtils._build_container_id( + managed_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="azure-router-model", container_id="cfile_upstream_abc", @@ -1271,7 +1271,7 @@ async def test_init_containers_api_endpoints_managed_id_without_model_id_unwraps router = Router(model_list=[]) mock_original_function = AsyncMock(return_value={"ok": True}) - managed_id = ResponsesAPIRequestUtils._build_container_id( + managed_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="openai", model_id=None, container_id="cfile_upstream_abc", @@ -1304,7 +1304,7 @@ async def test_init_containers_api_endpoints_managed_id_without_model_id_applies router = Router(model_list=[]) mock_original_function = AsyncMock(return_value={"ok": True}) - managed_id = ResponsesAPIRequestUtils._build_container_id( + managed_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id=None, container_id="cfile_upstream_abc", diff --git a/tests/router_unit_tests/test_router_handle_error.py b/tests/router_unit_tests/test_router_handle_error.py index 6b57efc7f37..aff133d1e74 100644 --- a/tests/router_unit_tests/test_router_handle_error.py +++ b/tests/router_unit_tests/test_router_handle_error.py @@ -136,7 +136,7 @@ async def test_async_raise_no_deployment_exception(): ] with patch( - "litellm.router_utils.handle_error._async_get_cooldown_deployments_with_debug_info", + "litellm.router_utils.handle_error.async_get_cooldown_deployments_with_debug_info", return_value=mock_cooldown_list, ): # Call the function @@ -188,7 +188,7 @@ async def test_async_raise_no_deployment_exception_empty_cooldown_list(): mock_cooldown_list: List = [] with patch( - "litellm.router_utils.handle_error._async_get_cooldown_deployments_with_debug_info", + "litellm.router_utils.handle_error.async_get_cooldown_deployments_with_debug_info", return_value=mock_cooldown_list, ): # Call the function @@ -232,7 +232,7 @@ async def test_async_raise_no_deployment_exception_none_cooldown_list(): mock_cooldown_list = None with patch( - "litellm.router_utils.handle_error._async_get_cooldown_deployments_with_debug_info", + "litellm.router_utils.handle_error.async_get_cooldown_deployments_with_debug_info", return_value=mock_cooldown_list, ): # After the defensive fix, this should handle None gracefully and return empty list diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index d3ad1d989c8..bad7ebdeb86 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -613,7 +613,7 @@ def test_deployment_callback_respects_cooldown_time(model_list): }, } - with patch("litellm.router._set_cooldown_deployments") as mock_set: + with patch("litellm.router.set_cooldown_deployments") as mock_set: router.deployment_callback_on_failure( kwargs=kwargs, completion_response=None, @@ -1565,7 +1565,7 @@ async def test_set_response_headers_wraps_bare_async_generator(model_list): def test_get_all_deployments(model_list): """Test if the 'get_all_deployments' function is working correctly""" router = Router(model_list=model_list) - deployments = router._get_all_deployments( + deployments = router.get_all_deployments( model_name="gpt-5-mini", model_alias="gpt-5-mini" ) assert len(deployments) > 0 @@ -1604,10 +1604,10 @@ def test_filter_cooldown_deployments(model_list): """Test if the 'filter_cooldown_deployments' function is working correctly""" router = Router(model_list=model_list) deployments = router._filter_cooldown_deployments( - healthy_deployments=router._get_all_deployments(model_name="gpt-5-mini"), # type: ignore + healthy_deployments=router.get_all_deployments(model_name="gpt-5-mini"), # type: ignore cooldown_deployments=[], ) - assert len(deployments) == len(router._get_all_deployments(model_name="gpt-5-mini")) + assert len(deployments) == len(router.get_all_deployments(model_name="gpt-5-mini")) def test_track_deployment_metrics(model_list): @@ -1780,7 +1780,7 @@ def test_get_model_from_alias(model_list): model_list=model_list, model_group_alias={"gpt-5.5": "gpt-5-mini"}, ) - model = router._get_model_from_alias(model="gpt-5.5") + model = router.get_model_from_alias(model="gpt-5.5") assert model == "gpt-5-mini" @@ -1869,7 +1869,7 @@ def test_pattern_match_deployment_set_model_name( import re # Convert model_name into a proper regex - model_name_regex = pattern_router._pattern_to_regex(model_name) + model_name_regex = pattern_router.pattern_to_regex(model_name) # Match against the request match = re.match(model_name_regex, user_request_model) @@ -2412,12 +2412,12 @@ def test_handle_clientside_credential_metadata_variable_name( model_list, function_name, metadata_key ): """Test that _handle_clientside_credential uses the correct metadata variable name based on function name""" - from litellm.router_utils.batch_utils import _get_router_metadata_variable_name + from litellm.router_utils.batch_utils import get_router_metadata_variable_name router = Router(model_list=model_list) # Verify the metadata variable name is correct for each function - expected_metadata_key = _get_router_metadata_variable_name( + expected_metadata_key = get_router_metadata_variable_name( function_name=function_name ) assert expected_metadata_key == metadata_key @@ -3095,7 +3095,7 @@ def test_sync_deployment_budget_config(monkeypatch): router._sync_deployment_budget_config(deployment=deployment) - budget_limiter = router._get_router_deployment_budget_limiter() + budget_limiter = router.get_router_deployment_budget_limiter() assert budget_limiter is not None config = budget_limiter._get_budget_config_for_deployment( "runtime-budget-deployment" @@ -3131,7 +3131,7 @@ def test_sync_deployment_budget_config_clears_removed_limits(monkeypatch): ) router._sync_deployment_budget_config(deployment=budgeted) - budget_limiter = router._get_router_deployment_budget_limiter() + budget_limiter = router.get_router_deployment_budget_limiter() assert budget_limiter is not None assert budget_limiter._get_budget_config_for_deployment(model_id) is not None @@ -3166,7 +3166,7 @@ def test_upsert_deployment_clears_stale_budget_config(monkeypatch): ) router.upsert_deployment(deployment=budgeted) - budget_limiter = router._get_router_deployment_budget_limiter() + budget_limiter = router.get_router_deployment_budget_limiter() assert budget_limiter is not None assert budget_limiter._get_budget_config_for_deployment(model_id) is not None diff --git a/tests/search_tests/base_search_unit_tests.py b/tests/search_tests/base_search_unit_tests.py index 42a4927e7c6..7028f58a1a3 100644 --- a/tests/search_tests/base_search_unit_tests.py +++ b/tests/search_tests/base_search_unit_tests.py @@ -41,7 +41,7 @@ class BaseSearchTest(ABC): """ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") - litellm._turn_on_debug() + litellm.turn_on_debug() search_provider = self.get_search_provider() print("Search Provider=", search_provider) diff --git a/tests/search_tests/test_duckduckgo_search.py b/tests/search_tests/test_duckduckgo_search.py index 69d19edded7..6f0ded3b146 100644 --- a/tests/search_tests/test_duckduckgo_search.py +++ b/tests/search_tests/test_duckduckgo_search.py @@ -29,7 +29,7 @@ class TestDuckDuckGoSearch(BaseSearchTest): """ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") - litellm._turn_on_debug() + litellm.turn_on_debug() search_provider = self.get_search_provider() print("Search Provider=", search_provider) diff --git a/tests/search_tests/test_perplexity_search.py b/tests/search_tests/test_perplexity_search.py index e1189a71355..1af0c18bc3e 100644 --- a/tests/search_tests/test_perplexity_search.py +++ b/tests/search_tests/test_perplexity_search.py @@ -34,7 +34,7 @@ class TestRouterSearch: from litellm import Router import litellm - litellm._turn_on_debug() + litellm.turn_on_debug() # Create router with search_tools config router = Router( diff --git a/tests/test_litellm/__init__.py b/tests/test_litellm/__init__.py index a2f3cf8de0f..7f5464af889 100644 --- a/tests/test_litellm/__init__.py +++ b/tests/test_litellm/__init__.py @@ -1 +1 @@ -# This file makes the tests/litellm directory a Python package +# This file makes the tests/litellm directory a Python package diff --git a/tests/test_litellm_rust/conftest.py b/tests/test_litellm_rust/conftest.py index a9ff759f0cf..2e965fe2630 100644 --- a/tests/test_litellm_rust/conftest.py +++ b/tests/test_litellm_rust/conftest.py @@ -12,7 +12,7 @@ import litellm from litellm import utils from litellm.litellm_core_utils import litellm_logging, thread_pool_executor from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER -from litellm.rust_bridge.configuration import ( # pyright: ignore[reportPrivateUsage] # preserve raw configuration state in test isolation +from litellm.rust_bridge.configuration import ( _CONFIGURATION, _parse_env_bool, ) diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index 264a666c685..18c8861ed6d 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -96,7 +96,7 @@ async def test_metadata_failure_dispatches_only_failure_and_releases_logger( seen: Final = [] class FailingMetadata(Logging): - def _response_cost_calculator(self, *args, **kwargs): + def response_cost_calculator(self, *args, **kwargs): raise failure def success_handler(self, *args, **kwargs): @@ -512,7 +512,7 @@ def created_loggers(monkeypatch: pytest.MonkeyPatch) -> list[Logging]: ) -> tuple[Logging, dict[str, object]]: logger, prepared = original_setup(call_type, rules, start, *args, is_async_call=is_async_call, **kwargs) assert isinstance(logger, Logging) - setattr(logger, "_defer_async_logging", True) + setattr(logger, "defer_async_logging", True) loggers.append(logger) return logger, prepared diff --git a/tests/test_litellm_rust/ocr/test_passthrough.py b/tests/test_litellm_rust/ocr/test_passthrough.py index b90979392b4..92493cb65ba 100644 --- a/tests/test_litellm_rust/ocr/test_passthrough.py +++ b/tests/test_litellm_rust/ocr/test_passthrough.py @@ -69,7 +69,7 @@ def test_ocr_relay_is_costed_per_page(model: str, native_path: str, body: object assert result.usage_info.pages_processed == pages assert logging_obj.call_type == "aocr" assert per_page is not None and per_page > 0 - assert logging_obj._response_cost_calculator(result=result) == pytest.approx(pages * per_page) # pyright: ignore[reportPrivateUsage] # the per-page cost path is what the relay routes into + assert logging_obj.response_cost_calculator(result=result) == pytest.approx(pages * per_page) @pytest.mark.parametrize( diff --git a/tests/test_litellm_rust/support/isolation.py b/tests/test_litellm_rust/support/isolation.py index 26c7cd0f875..4464ee4bb51 100644 --- a/tests/test_litellm_rust/support/isolation.py +++ b/tests/test_litellm_rust/support/isolation.py @@ -53,6 +53,6 @@ def isolated_callback_registries() -> Generator[None]: with ExitStack() as stack: for attribute in CALLBACK_ATTRIBUTES: stack.enter_context(_isolated_list(litellm, attribute)) - stack.enter_context(_isolated_list(litellm_logging, "_in_memory_loggers")) # pyright: ignore[reportPrivateUsage] # no public callback-cache accessor + stack.enter_context(_isolated_list(litellm_logging, "_in_memory_loggers")) stack.enter_context(rebound(utils, "callback_list", [])) yield diff --git a/tests/unified_google_tests/base_google_test.py b/tests/unified_google_tests/base_google_test.py index d6de60f6ec2..558ec2445fb 100644 --- a/tests/unified_google_tests/base_google_test.py +++ b/tests/unified_google_tests/base_google_test.py @@ -173,7 +173,7 @@ class BaseGoogleGenAITest: if temp_file_path: self._temp_files_to_cleanup.append(temp_file_path) - litellm._turn_on_debug() + litellm.turn_on_debug() print( f"Testing {'async' if is_async else 'sync'} non-streaming with model config: {request_params}" @@ -197,7 +197,7 @@ class BaseGoogleGenAITest: @pytest.mark.asyncio async def test_async_non_streaming_with_logging(self): """Test async non-streaming Google GenAI generate content with logging""" - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.logging_callback_manager._reset_all_callbacks() litellm.set_verbose = True test_custom_logger = TestCustomLogger() @@ -236,7 +236,7 @@ class BaseGoogleGenAITest: @pytest.mark.asyncio async def test_async_streaming_with_logging(self): """Test async streaming Google GenAI generate content with logging""" - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True litellm.logging_callback_manager._reset_all_callbacks() test_custom_logger = TestCustomLogger() diff --git a/tests/unified_google_tests/base_interactions_test.py b/tests/unified_google_tests/base_interactions_test.py index 6386193e580..4dff91e697e 100644 --- a/tests/unified_google_tests/base_interactions_test.py +++ b/tests/unified_google_tests/base_interactions_test.py @@ -32,7 +32,7 @@ class BaseInteractionsTest(ABC): def test_create_simple_string_input(self): """Test creating an interaction with a simple string input.""" - litellm._turn_on_debug() + litellm.turn_on_debug() api_key = self.get_api_key() if not api_key: pytest.skip(f"API key not set for {self.__class__.__name__}") diff --git a/tests/unified_google_tests/test_google_ai_studio.py b/tests/unified_google_tests/test_google_ai_studio.py index 3c4213bbb29..0d9f8365ecf 100644 --- a/tests/unified_google_tests/test_google_ai_studio.py +++ b/tests/unified_google_tests/test_google_ai_studio.py @@ -28,7 +28,7 @@ async def test_mock_stream_generate_content_with_tools(): """Test streaming function call response parsing and validation""" from litellm.types.google_genai.main import ToolConfigDict - litellm._turn_on_debug() + litellm.turn_on_debug() contents = [ { "role": "user", diff --git a/tests/unit/batches/test_batch_utils.py b/tests/unit/batches/test_batch_utils.py index 8147626fc5e..52217251b6c 100644 --- a/tests/unit/batches/test_batch_utils.py +++ b/tests/unit/batches/test_batch_utils.py @@ -86,6 +86,13 @@ def test_get_response_body_present(): } +def test_get_response_body_is_returned_without_validation(): + response_body = ["provider-specific response"] + row = {"response": {"body": response_body}} + + assert bu._get_response_from_batch_job_output_file(row) is response_body + + @pytest.mark.parametrize( "row", [ @@ -347,6 +354,10 @@ def test_count_entry_messages_path(fake_token_counter): assert bu._count_entry_tokens(entry) == 2 # len(messages) +def test_count_entry_uses_dynamic_length_for_messages(fake_token_counter): + assert bu._count_entry_tokens({"body": {"messages": "abc"}}) == 3 + + def test_count_entry_prompt_path(fake_token_counter): assert bu._count_entry_tokens({"body": {"model": "gpt-4o", "prompt": "abcd"}}) == 4 @@ -1895,9 +1906,9 @@ class TestFileAccessCredentialsCarryFederation: has to inherit the federation fields or it cannot authenticate and the batch is never billed.""" def test_federation_fields_survive_extraction(self): - from litellm.batches.batch_utils import _extract_file_access_credentials + from litellm.batches.batch_utils import extract_file_access_credentials - credentials = _extract_file_access_credentials( + credentials = extract_file_access_credentials( { "model": "anthropic/claude-sonnet-4-5", "anthropic_federation_rule_id": "fdrl_x", @@ -1914,12 +1925,12 @@ class TestFileAccessCredentialsCarryFederation: def test_every_federation_field_is_carried(self): """Derived from the kwargs set, so a new federation field is carried without an edit here.""" - from litellm.batches.batch_utils import _extract_file_access_credentials + from litellm.batches.batch_utils import extract_file_access_credentials from litellm.litellm_core_utils.get_litellm_params import ANTHROPIC_WIF_KWARGS_KEYS params = {name: f"value-{name}" for name in ANTHROPIC_WIF_KWARGS_KEYS} - credentials = _extract_file_access_credentials(params) + credentials = extract_file_access_credentials(params) assert set(credentials) == set(ANTHROPIC_WIF_KWARGS_KEYS) diff --git a/tests/unit/caching/test_caching_handler.py b/tests/unit/caching/test_caching_handler.py index abbd44384b1..29398a45c95 100644 --- a/tests/unit/caching/test_caching_handler.py +++ b/tests/unit/caching/test_caching_handler.py @@ -123,7 +123,7 @@ async def test_async_set_get_cache(response): await asyncio.sleep(2) # Verify the result was cached - cached_response = await caching_handler._async_get_cache( + cached_response = await caching_handler.async_get_cache( model="gpt-3.5-turbo", original_function=original_function, logging_obj=logging_obj, @@ -285,7 +285,7 @@ def test_combine_cached_embedding_response_with_api_result(): ) # Call the method - result = caching_handler._combine_cached_embedding_response_with_api_result( + result = caching_handler.combine_cached_embedding_response_with_api_result( _caching_handler_response=caching_handler_response, embedding_response=api_response, start_time=start_time, @@ -342,7 +342,7 @@ def test_combine_cached_embedding_response_multiple_missing_values(): ) # Call the method - result = caching_handler._combine_cached_embedding_response_with_api_result( + result = caching_handler.combine_cached_embedding_response_with_api_result( _caching_handler_response=caching_handler_response, embedding_response=api_response, start_time=start_time, @@ -405,7 +405,7 @@ async def test_embedding_cache_model_field_consistency(): ) # Step 2: Retrieve from cache - cached_response = await caching_handler._async_get_cache( + cached_response = await caching_handler.async_get_cache( model=original_model, original_function=aembedding, logging_obj=logging_obj, @@ -487,7 +487,7 @@ async def test_embedding_cache_model_field_with_vendor_prefix(): ) # Retrieve from cache - cached_response = await caching_handler._async_get_cache( + cached_response = await caching_handler.async_get_cache( model=vendor_model, original_function=aembedding, logging_obj=logging_obj, @@ -620,7 +620,7 @@ async def test_async_responses_api_caching(): await asyncio.sleep(0.5) # Step 2: Retrieve from cache - cached_response = await caching_handler._async_get_cache( + cached_response = await caching_handler.async_get_cache( model=original_model, original_function=aresponses, logging_obj=logging_obj, @@ -672,7 +672,7 @@ async def test_async_get_cache_updates_request_kwargs_for_streaming_responses(): "caching": True, } - await caching_handler._async_get_cache( + await caching_handler.async_get_cache( model="gpt-4o", original_function=aresponses, logging_obj=logging_obj, @@ -749,7 +749,7 @@ def test_sync_responses_api_caching(): caching_handler.sync_set_cache(result=responses_api_response, kwargs=kwargs) # Step 2: Retrieve from cache - cached_response = caching_handler._sync_get_cache( + cached_response = caching_handler.sync_get_cache( model=original_model, original_function=responses, logging_obj=logging_obj, @@ -881,7 +881,7 @@ def test_sync_get_cache_does_not_eagerly_log_streaming_responses_hits(): caching_handler.sync_set_cache(result=responses_api_response, kwargs=kwargs) - cached_response = caching_handler._sync_get_cache( + cached_response = caching_handler.sync_get_cache( model=original_model, original_function=responses, logging_obj=logging_obj, @@ -925,7 +925,7 @@ def test_sync_get_cache_defers_streaming_completion_hit_callbacks(): caching_handler.sync_set_cache(result=chat_completion_response, kwargs=kwargs) - cached_response = caching_handler._sync_get_cache( + cached_response = caching_handler.sync_get_cache( model=original_model, original_function=completion, logging_obj=logging_obj, @@ -983,7 +983,7 @@ async def test_async_get_cache_defers_streaming_completion_hit_callbacks(): ) caching_handler._async_log_cache_hit_on_callbacks = MagicMock() - cached_response = await caching_handler._async_get_cache( + cached_response = await caching_handler.async_get_cache( model=original_model, original_function=litellm.acompletion, logging_obj=logging_obj, @@ -1312,7 +1312,7 @@ async def test_responses_api_cache_with_different_inputs(): start_time=datetime.now(), ) - cached_1 = await caching_handler._async_get_cache( + cached_1 = await caching_handler.async_get_cache( model=original_model, original_function=aresponses, logging_obj=logging_obj_1, @@ -1321,7 +1321,7 @@ async def test_responses_api_cache_with_different_inputs(): kwargs=kwargs_1, ) - cached_2 = await caching_handler._async_get_cache( + cached_2 = await caching_handler.async_get_cache( model=original_model, original_function=aresponses, logging_obj=logging_obj_2, @@ -1996,7 +1996,7 @@ def test_sync_stream_responses_cache_hit_sets_custom_llm_provider_on_logging_obj handler = LLMCachingHandler(original_function=litellm.responses, request_kwargs=kwargs, start_time=datetime.now()) logging_obj = _build_logging_obj(CallTypes.responses.value, stream=True) - hit = handler._sync_get_cache( + hit = handler.sync_get_cache( model="azure/gpt-5.4-mini", original_function=litellm.responses, logging_obj=logging_obj, @@ -2013,7 +2013,7 @@ def test_sync_stream_responses_cache_hit_sets_custom_llm_provider_on_logging_obj def test_request_kwargs_does_not_retain_logging_obj(): """ - The caching handler lives on logging_obj._llm_caching_handler, so keeping + The caching handler lives on logging_obj.llm_caching_handler, so keeping litellm_logging_obj inside request_kwargs closes a reference cycle (Logging -> LLMCachingHandler -> kwargs -> Logging). That cycle keeps the full request payload alive until a generational GC pass instead of being @@ -2130,7 +2130,7 @@ async def test_cache_hit_records_the_looked_up_key_as_the_preset_cache_key(monke logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False) logging_obj.async_success_handler = AsyncMock() - hit = await handler._async_get_cache( + hit = await handler.async_get_cache( model="gpt-5.4", original_function=acompletion, logging_obj=logging_obj, @@ -2176,7 +2176,7 @@ async def test_converted_stream_cache_hit_replayed_as_plain_object_logs_at_hit_t logging_obj.async_success_handler = AsyncMock() logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() - hit = await handler._async_get_cache( + hit = await handler.async_get_cache( model="claude-sonnet-5", original_function=aanthropic_messages, logging_obj=logging_obj, @@ -2217,7 +2217,7 @@ async def test_agentic_loop_followup_cache_hit_with_converted_stream_marker_repl logging_obj.async_success_handler = AsyncMock() logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() - hit = await handler._async_get_cache( + hit = await handler.async_get_cache( model="gpt-5.6", original_function=acompletion, logging_obj=logging_obj, @@ -2290,7 +2290,7 @@ async def test_response_cache_lookup_and_write_declare_the_llm_response_target(m def get_cache_key(self, **kwargs): return "k" - def _supports_async(self): + def supports_async(self): return True async def async_get_cache(self, **kwargs): @@ -2306,7 +2306,7 @@ async def test_response_cache_lookup_and_write_declare_the_llm_response_target(m handler = LLMCachingHandler(original_function=acompletion, request_kwargs={}, start_time=datetime.now()) monkeypatch.setattr(litellm, "cache", _TargetRecordingCache()) - await handler._async_get_cache( + await handler.async_get_cache( model="gpt-3.5-turbo", original_function=acompletion, logging_obj=MagicMock(), @@ -2457,7 +2457,7 @@ async def test_async_set_cache_skips_response_without_output( await asyncio.gather(*_PENDING_CACHE_WRITES) assert await litellm.cache.async_get_cache(**kwargs) is None - lookup = await handler._async_get_cache( + lookup = await handler.async_get_cache( model="gpt-3.5-turbo", original_function=original_function, logging_obj=_completion_logging_obj(call_type), @@ -2482,7 +2482,7 @@ async def test_async_get_cache_treats_stored_response_without_output_as_miss( assert await litellm.cache.async_get_cache(**kwargs) is not None handler = LLMCachingHandler(original_function=original_function, request_kwargs={}, start_time=_FIXED_START) - lookup = await handler._async_get_cache( + lookup = await handler.async_get_cache( model="gpt-3.5-turbo", original_function=original_function, logging_obj=_completion_logging_obj(call_type), @@ -2505,7 +2505,7 @@ async def test_async_get_cache_heals_stored_completion_without_choices(): assert await litellm.cache.async_get_cache(**kwargs) is not None async def lookup(): - return await handler._async_get_cache( + return await handler.async_get_cache( model="gpt-3.5-turbo", original_function=litellm.acompletion, logging_obj=_completion_logging_obj(CallTypes.acompletion.value), @@ -2531,7 +2531,7 @@ def test_sync_set_cache_skips_response_without_output(): handler.sync_set_cache(result=litellm.ModelResponse(choices=[]), kwargs=kwargs) assert litellm.cache.get_cache(**kwargs) is None - lookup = handler._sync_get_cache( + lookup = handler.sync_get_cache( model="gpt-3.5-turbo", original_function=completion, logging_obj=_completion_logging_obj(CallTypes.completion.value), @@ -2550,7 +2550,7 @@ def test_sync_get_cache_heals_stored_completion_without_choices(): assert litellm.cache.get_cache(**kwargs) is not None def lookup(): - return handler._sync_get_cache( + return handler.sync_get_cache( model="gpt-3.5-turbo", original_function=completion, logging_obj=_completion_logging_obj(CallTypes.completion.value), @@ -2609,7 +2609,7 @@ async def test_async_streamed_answer_without_output_is_never_cached(): kwargs = {"model": "gpt-4o", "messages": _unique_messages()} handler = LLMCachingHandler(original_function=litellm.acompletion, request_kwargs=kwargs, start_time=_FIXED_START) - await handler._add_streaming_response_to_cache(_closing_chunk_without_output()) + await handler.add_streaming_response_to_cache(_closing_chunk_without_output()) await asyncio.gather(*_PENDING_CACHE_WRITES) assert litellm.cache.get_cache(**kwargs) is None @@ -2621,8 +2621,8 @@ async def test_async_streamed_answer_with_content_is_cached(): kwargs = {"model": "gpt-4o", "messages": _unique_messages()} handler = LLMCachingHandler(original_function=litellm.acompletion, request_kwargs=kwargs, start_time=_FIXED_START) - await handler._add_streaming_response_to_cache(_content_chunk("hi")) - await handler._add_streaming_response_to_cache(_closing_chunk_without_output()) + await handler.add_streaming_response_to_cache(_content_chunk("hi")) + await handler.add_streaming_response_to_cache(_closing_chunk_without_output()) await asyncio.gather(*_PENDING_CACHE_WRITES) assert _cached_content(litellm.cache.get_cache(**kwargs)) == "hi" @@ -2633,7 +2633,7 @@ def test_sync_streamed_answer_without_output_is_never_cached(): kwargs = {"model": "gpt-4o", "messages": _unique_messages()} handler = LLMCachingHandler(original_function=completion, request_kwargs=kwargs, start_time=_FIXED_START) - handler._sync_add_streaming_response_to_cache(_closing_chunk_without_output()) + handler.sync_add_streaming_response_to_cache(_closing_chunk_without_output()) assert litellm.cache.get_cache(**kwargs) is None @@ -2643,8 +2643,8 @@ def test_sync_streamed_answer_with_content_is_cached(): kwargs = {"model": "gpt-4o", "messages": _unique_messages()} handler = LLMCachingHandler(original_function=completion, request_kwargs=kwargs, start_time=_FIXED_START) - handler._sync_add_streaming_response_to_cache(_content_chunk("hi")) - handler._sync_add_streaming_response_to_cache(_closing_chunk_without_output()) + handler.sync_add_streaming_response_to_cache(_content_chunk("hi")) + handler.sync_add_streaming_response_to_cache(_closing_chunk_without_output()) assert _cached_content(litellm.cache.get_cache(**kwargs)) == "hi" @@ -2660,7 +2660,7 @@ async def test_async_get_cache_forgets_the_worker_copy_of_a_stored_response_with await handler.dual_cache.async_set_cache(key, {"timestamp": _FIXED_START.timestamp(), "response": poisoned}) assert await handler.dual_cache.async_get_cache(key) is not None - lookup = await handler._async_get_cache( + lookup = await handler.async_get_cache( model="gpt-3.5-turbo", original_function=litellm.acompletion, logging_obj=_completion_logging_obj(CallTypes.acompletion.value), diff --git a/tests/unit/caching/test_redis_semantic_cache.py b/tests/unit/caching/test_redis_semantic_cache.py index 461689165bb..99844c695cb 100644 --- a/tests/unit/caching/test_redis_semantic_cache.py +++ b/tests/unit/caching/test_redis_semantic_cache.py @@ -906,7 +906,7 @@ def test_cache_get_cache_filters_non_lookup_kwargs_from_backend_cache(): cache.cache = MagicMock() cache.should_use_cache = MagicMock(return_value=True) cache.get_cache_key = MagicMock(return_value="test_key") - cache._get_cache_logic = MagicMock(return_value={"content": "Paris"}) + cache.get_cache_logic = MagicMock(return_value={"content": "Paris"}) def _cache_hit(_cache_key, **cache_kwargs): cache_kwargs["metadata"]["semantic-similarity"] = 0.7 @@ -940,7 +940,7 @@ def test_cache_get_cache_filters_non_lookup_kwargs_from_backend_cache(): }, } assert forwarded_kwargs["metadata"] is not metadata - cache._get_cache_logic.assert_called_once_with( + cache.get_cache_logic.assert_called_once_with( cached_result={"content": "Paris"}, max_age=10, ) @@ -954,7 +954,7 @@ def test_cache_get_cache_filters_sensitive_kwargs_without_metadata(): cache.cache.get_cache = MagicMock(return_value={"content": "Paris"}) cache.should_use_cache = MagicMock(return_value=True) cache.get_cache_key = MagicMock(return_value="test_key") - cache._get_cache_logic = MagicMock(return_value={"content": "Paris"}) + cache.get_cache_logic = MagicMock(return_value={"content": "Paris"}) result = cache.get_cache( input="What is the capital of France?", @@ -976,7 +976,7 @@ def test_cache_get_cache_passes_responses_input_to_dynamic_cache(): cache = Cache.__new__(Cache) cache.should_use_cache = MagicMock(return_value=True) cache.get_cache_key = MagicMock(return_value="test_key") - cache._get_cache_logic = MagicMock(return_value={"content": "Paris"}) + cache.get_cache_logic = MagicMock(return_value={"content": "Paris"}) dynamic_cache_object = MagicMock() dynamic_cache_object.get_cache = MagicMock(return_value={"content": "Paris"}) @@ -994,7 +994,7 @@ def test_cache_get_cache_passes_responses_input_to_dynamic_cache(): input="What is the capital of France?", metadata=metadata, ) - cache._get_cache_logic.assert_called_once_with( + cache.get_cache_logic.assert_called_once_with( cached_result={"content": "Paris"}, max_age=float("inf"), ) diff --git a/tests/unit/caching/test_request_redis_batch_pre_call.py b/tests/unit/caching/test_request_redis_batch_pre_call.py index cee7bfd8c65..81a2a1718e9 100644 --- a/tests/unit/caching/test_request_redis_batch_pre_call.py +++ b/tests/unit/caching/test_request_redis_batch_pre_call.py @@ -283,7 +283,7 @@ def _deployment(deployment_id: str) -> dict: def _router(redis_cache: FakeRedisCache, routing_strategy: str = "usage-based-routing-v2") -> Router: router = Router(model_list=[_deployment("dep-a"), _deployment("dep-b")], routing_strategy=routing_strategy) - router._update_redis_cache(cache=redis_cache) + router.update_redis_cache(cache=redis_cache) return router @@ -689,7 +689,7 @@ async def test_simple_shuffle_prefetches_only_its_cooldown_read_into_the_admissi assert redis_cache.alone == [] shuffle = Router(model_list=[_deployment("dep-a")], routing_strategy="simple-shuffle") - shuffle._update_redis_cache(cache=redis_cache) + shuffle.update_redis_cache(cache=redis_cache) with request_redis_batch_scope() as request: shuffle.arm_routing_read_prefetch(_MODEL_GROUP, {}) armed = request.prefetched["routing_read"] diff --git a/tests/unit/caching/test_responses_stream_cache_keys.py b/tests/unit/caching/test_responses_stream_cache_keys.py index 5637028f550..e836948f05e 100644 --- a/tests/unit/caching/test_responses_stream_cache_keys.py +++ b/tests/unit/caching/test_responses_stream_cache_keys.py @@ -32,7 +32,7 @@ async def test_async_get_cache_reuses_preset_cache_key_for_responses(): original_cache = litellm.cache mock_cache = MagicMock() mock_cache.supported_call_types = [CallTypes.aresponses.value] - mock_cache._supports_async.return_value = True + mock_cache.supports_async.return_value = True mock_cache.get_cache_key.return_value = "responses-stream-cache-key" mock_cache.async_get_cache = AsyncMock(return_value=None) litellm.cache = mock_cache @@ -43,7 +43,7 @@ async def test_async_get_cache_reuses_preset_cache_key_for_responses(): "stream": True, "litellm_params": {}, } - await caching_handler._async_get_cache( + await caching_handler.async_get_cache( model="gpt-4.1-mini", original_function=aresponses, logging_obj=logging_obj, @@ -82,7 +82,7 @@ async def test_async_get_cache_falls_back_to_sync_cache_for_responses(): original_cache = litellm.cache mock_cache = MagicMock() mock_cache.supported_call_types = [CallTypes.aresponses.value] - mock_cache._supports_async.return_value = False + mock_cache.supports_async.return_value = False mock_cache.get_cache_key.return_value = "responses-stream-cache-key" mock_cache.get_cache.return_value = None litellm.cache = mock_cache @@ -93,7 +93,7 @@ async def test_async_get_cache_falls_back_to_sync_cache_for_responses(): "stream": True, "litellm_params": {}, } - await caching_handler._async_get_cache( + await caching_handler.async_get_cache( model="gpt-4.1-mini", original_function=aresponses, logging_obj=logging_obj, diff --git a/tests/unit/caching/test_s3_cache.py b/tests/unit/caching/test_s3_cache.py index f86f2da30ef..4ea98322f13 100644 --- a/tests/unit/caching/test_s3_cache.py +++ b/tests/unit/caching/test_s3_cache.py @@ -331,7 +331,7 @@ def test_s3_cache_supports_async(): cache = Cache(type=LiteLLMCacheType.S3, s3_bucket_name="test-bucket") # Should now return True for async support - assert cache._supports_async() is True + assert cache.supports_async() is True @pytest.mark.asyncio diff --git a/tests/unit/caching/test_unit_test_caching.py b/tests/unit/caching/test_unit_test_caching.py index d720838c507..83132b5aa1d 100644 --- a/tests/unit/caching/test_unit_test_caching.py +++ b/tests/unit/caching/test_unit_test_caching.py @@ -36,7 +36,7 @@ import logging def test_get_kwargs_for_cache_key(): _cache = litellm.Cache() - relevant_kwargs = ModelParamHelper._get_all_llm_api_params() + relevant_kwargs = ModelParamHelper.get_all_llm_api_params() print(relevant_kwargs) diff --git a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 25a3220792f..4d89f021fe9 100644 --- a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -1577,7 +1577,7 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): del os.environ["LITELLM_REASONING_AUTO_SUMMARY"] for effort in effort_levels: - result = handler._map_reasoning_effort(effort) + result = handler.map_reasoning_effort(effort) assert result is not None, f"Result should not be None for effort={effort}" assert result["effort"] == effort, f"Effort should be {effort}" @@ -1593,7 +1593,7 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): litellm.reasoning_auto_summary = True for effort in effort_levels: - result = handler._map_reasoning_effort(effort) + result = handler.map_reasoning_effort(effort) assert result is not None, f"Result should not be None for effort={effort}" assert result["effort"] == effort, f"Effort should be {effort}" @@ -1609,7 +1609,7 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): litellm.reasoning_auto_summary = False monkeypatch.setenv("LITELLM_REASONING_AUTO_SUMMARY", "true") - result = handler._map_reasoning_effort("high") + result = handler.map_reasoning_effort("high") assert ( result["summary"] == "detailed" ), "Summary should be 'detailed' when env var is enabled" @@ -1621,7 +1621,7 @@ def test_map_reasoning_effort_adds_summary_detailed(monkeypatch): del os.environ["LITELLM_REASONING_AUTO_SUMMARY"] dict_input = {"effort": "high", "summary": "custom_summary"} - result_dict = handler._map_reasoning_effort(dict_input) + result_dict = handler.map_reasoning_effort(dict_input) assert result_dict["effort"] == "high" assert result_dict["summary"] == "custom_summary" print("✓ Dict input is passed through without modification") diff --git a/tests/unit/containers/test_azure_container_transformation.py b/tests/unit/containers/test_azure_container_transformation.py index 1c990220e11..38ea9323017 100644 --- a/tests/unit/containers/test_azure_container_transformation.py +++ b/tests/unit/containers/test_azure_container_transformation.py @@ -658,7 +658,7 @@ class TestAzureContainerKnownFailureRegressions: from litellm.proxy.container_endpoints import handler_factory - encoded_id = ResponsesAPIRequestUtils._build_container_id( + encoded_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="model_abc123", container_id="cntr_123", @@ -729,7 +729,7 @@ class TestAzureContainerKnownFailureRegressions: from litellm.proxy.container_endpoints import handler_factory - encoded_id = ResponsesAPIRequestUtils._build_container_id( + encoded_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="model_abc123", container_id="cntr_123", @@ -807,7 +807,7 @@ class TestAzureContainerKnownFailureRegressions: from litellm.proxy.common_utils import http_parsing_utils from litellm.proxy.container_endpoints import handler_factory - encoded_id = ResponsesAPIRequestUtils._build_container_id( + encoded_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="model_abc123", container_id="cntr_123", @@ -893,7 +893,7 @@ class TestAzureContainerKnownFailureRegressions: get_container_forwarding_params, ) - encoded_id = ResponsesAPIRequestUtils._build_container_id( + encoded_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="deployment-uuid-123", container_id="cntr_6a058b43d24c8190a226cfb1d35405b20115fb7875ff11df", @@ -932,7 +932,7 @@ class TestAzureContainerKnownFailureRegressions: ) native_id = "cntr_6a058b43d24c8190a226cfb1d35405b20115fb7875ff11df" - encoded_stored_id = ResponsesAPIRequestUtils._build_container_id( + encoded_stored_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="deployment-uuid-123", container_id=native_id, diff --git a/tests/unit/containers/test_container_api.py b/tests/unit/containers/test_container_api.py index 16a3844431e..edaf960db2c 100644 --- a/tests/unit/containers/test_container_api.py +++ b/tests/unit/containers/test_container_api.py @@ -183,7 +183,7 @@ class TestContainerAPI: litellm_metadata={"model_info": {"id": "deployment-abc"}}, ) - decoded = ResponsesAPIRequestUtils._decode_container_id(response.id) + decoded = ResponsesAPIRequestUtils.decode_container_id(response.id) assert decoded["model_id"] == "deployment-abc" assert decoded["custom_llm_provider"] == "openai" assert decoded["response_id"] == "cntr_upstream_123" @@ -272,7 +272,7 @@ class TestContainerAPI: def test_retrieve_container_reencodes_short_managed_id_for_routing(self): """Short cntr_ IDs must still re-encode output so follow-ups keep router affinity.""" - short_managed_id = ResponsesAPIRequestUtils._build_container_id( + short_managed_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="router-gpt", container_id="x", @@ -303,14 +303,14 @@ class TestContainerAPI: mock_method.assert_called_once() assert mock_method.call_args.kwargs["container_id"] == "x" assert response.id.startswith("cntr_") - decoded = ResponsesAPIRequestUtils._decode_container_id(response.id) + decoded = ResponsesAPIRequestUtils.decode_container_id(response.id) assert decoded.get("response_id") == "x" assert decoded.get("model_id") == "router-gpt" assert decoded.get("custom_llm_provider") == "azure" def test_delete_container_reencodes_short_managed_id_for_routing(self): """Same as retrieve: short managed IDs must round-trip encoding on delete result.""" - short_managed_id = ResponsesAPIRequestUtils._build_container_id( + short_managed_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="router-gpt", container_id="z", @@ -336,7 +336,7 @@ class TestContainerAPI: mock_method.assert_called_once() assert mock_method.call_args.kwargs["container_id"] == "z" assert response.id.startswith("cntr_") - decoded = ResponsesAPIRequestUtils._decode_container_id(response.id) + decoded = ResponsesAPIRequestUtils.decode_container_id(response.id) assert decoded.get("response_id") == "z" assert decoded.get("model_id") == "router-gpt" diff --git a/tests/unit/containers/test_container_proxy_ownership.py b/tests/unit/containers/test_container_proxy_ownership.py index 38988d65c04..dd21faa1059 100644 --- a/tests/unit/containers/test_container_proxy_ownership.py +++ b/tests/unit/containers/test_container_proxy_ownership.py @@ -449,7 +449,7 @@ async def test_should_validate_owner_and_forward_decoded_id_for_multipart_upload "convert_upload_files_to_file_data", AsyncMock(return_value={"file": ["file-data"]}), ) - encoded_id = ResponsesAPIRequestUtils._build_container_id( + encoded_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="router-gpt", container_id="cntr_provider", @@ -508,7 +508,7 @@ async def test_should_forward_decoded_container_id_for_proxy_retrieve(monkeypatc "assert_user_can_access_container", AsyncMock(return_value=("cntr_provider", "azure")), ) - encoded_id = ResponsesAPIRequestUtils._build_container_id( + encoded_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="router-gpt", container_id="cntr_provider", @@ -781,7 +781,7 @@ async def test_should_forward_decoded_container_id_for_proxy_delete(monkeypatch) "assert_user_can_access_container", AsyncMock(return_value=("cntr_provider", "azure")), ) - encoded_id = ResponsesAPIRequestUtils._build_container_id( + encoded_id = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id="router-gpt", container_id="cntr_provider", diff --git a/tests/unit/containers/test_container_utils.py b/tests/unit/containers/test_container_utils.py index a81d1263d6b..64ce8a7bf06 100644 --- a/tests/unit/containers/test_container_utils.py +++ b/tests/unit/containers/test_container_utils.py @@ -226,7 +226,7 @@ class TestContainerRequestUtils: def test_decode_managed_container_id_returns_provider_container_id(self): """Managed IDs must decode to the short ID sent on upstream requests.""" inner = "cntr_69d4ff00deadbeef" - managed = ResponsesAPIRequestUtils._build_container_id( + managed = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="openai", model_id=None, container_id=inner, diff --git a/tests/unit/enterprise/proxy/hooks/test_managed_files.py b/tests/unit/enterprise/proxy/hooks/test_managed_files.py index 7e96c956664..7d8c8b5c425 100644 --- a/tests/unit/enterprise/proxy/hooks/test_managed_files.py +++ b/tests/unit/enterprise/proxy/hooks/test_managed_files.py @@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException -from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles +from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles from litellm.caching import DualCache from litellm.proxy._types import CallTypes @@ -17,7 +17,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( def test_get_file_ids_from_messages(): - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) messages = [ @@ -41,7 +41,7 @@ def test_get_file_ids_from_messages(): def test_get_file_ids_from_messages_skips_bedrock_content_blocks_without_type(): - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) messages = [ @@ -81,7 +81,7 @@ async def test_async_pre_call_hook_batch_retrieve(): return_value = MagicMock() return_value.created_by = "123" prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) data = { @@ -107,7 +107,7 @@ async def test_list_user_batches_limit_zero_returns_empty_page_without_db_query( from litellm.proxy._types import UserAPIKeyAuth prisma_client = MagicMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache(), prisma_client=prisma_client) + proxy_managed_files = PROXY_LiteLLMManagedFiles(DualCache(), prisma_client=prisma_client) page = await proxy_managed_files.list_user_batches( user_api_key_dict=UserAPIKeyAuth(user_id="123"), @@ -126,7 +126,7 @@ async def test_async_pre_call_deployment_hook_resolves_model_id_from_litellm_met async_pre_call_deployment_hook must check both locations so the managed file ID is resolved to the provider-specific file ID. """ - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) @@ -160,7 +160,7 @@ async def test_async_pre_call_deployment_hook_prefers_top_level_model_info(): When model_info exists at top-level kwargs, async_pre_call_deployment_hook should use it without falling back to litellm_metadata. """ - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) @@ -199,7 +199,7 @@ async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_uncha When model_info is absent from both top-level and litellm_metadata, the managed file ID should remain unchanged. """ - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) @@ -222,7 +222,7 @@ async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_uncha # def test_list_managed_files(): -# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache()) +# proxy_managed_files = PROXY_LiteLLMManagedFiles(DualCache()) # # Create some test files # file1 = proxy_managed_files.create_file( @@ -245,7 +245,7 @@ async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_uncha # assert all(f.purpose == "assistants" for f in files) # def test_retrieve_managed_file(): -# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache()) +# proxy_managed_files = PROXY_LiteLLMManagedFiles(DualCache()) # # Create a test file # test_content = b"test content for retrieve" @@ -265,7 +265,7 @@ async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_uncha # assert retrieved_file.status == "uploaded" # def test_delete_managed_file(): -# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache()) +# proxy_managed_files = PROXY_LiteLLMManagedFiles(DualCache()) # # Create a test file # created_file = proxy_managed_files.create_file( @@ -289,21 +289,21 @@ async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_uncha # assert created_file.id not in [f.id for f in files] # def test_retrieve_nonexistent_file(): -# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache()) +# proxy_managed_files = PROXY_LiteLLMManagedFiles(DualCache()) # # Try to retrieve a non-existent file # with pytest.raises(Exception): # proxy_managed_files.retrieve_file("nonexistent-file-id") # def test_delete_nonexistent_file(): -# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache()) +# proxy_managed_files = PROXY_LiteLLMManagedFiles(DualCache()) # # Try to delete a non-existent file # with pytest.raises(Exception): # proxy_managed_files.delete_file("nonexistent-file-id") # def test_list_files_with_purpose_filter(): -# proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache()) +# proxy_managed_files = PROXY_LiteLLMManagedFiles(DualCache()) # # Create files with different purposes # file1 = proxy_managed_files.create_file( @@ -351,7 +351,7 @@ async def test_async_post_call_success_hook_for_unified_finetuning_job(): "unified_file_id": unified_file_id, "model_id": "gpt-3.5-turbo-0613", } - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=AsyncMock() ) data = { @@ -376,7 +376,7 @@ async def test_async_pre_call_hook_for_unified_finetuning_job(): return_value = MagicMock() return_value.created_by = "123" prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) data = { @@ -408,7 +408,7 @@ async def test_can_user_call_unified_file_id(call_type): return_value = MagicMock() return_value.created_by = "123" prisma_client.db.litellm_managedfiletable.find_first.return_value = return_value - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxmMTNlNDAzZS01YWM3LTRhZjktOGQzNS0wNDgwZDMxOTgyYTg7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00by1taW5pLW9wZW5haTtsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1Ib3UxZDFXc3c1SDNKcjFMYllpZDJiO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxmODBiNWU2NzQ1NzdkNjkyMjM4YmVhNTIxZDdiMGI5ZGYyY2FmMTEwMTU2YmU5YzBjM2NjMmNkNTBjOTM1ZDI0" @@ -435,7 +435,7 @@ async def test_router_acreate_batch_only_selects_from_file_id_mapping(monkeypatc return_value = MagicMock() return_value.created_by = "123" prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -522,7 +522,7 @@ async def test_output_file_id_for_batch_retrieve(): "litellm_model_name": "gpt-5.5", "unified_batch_id": "litellm_proxy;model_id:12345679;llm_batch_id:batch_685c5e5d63988190b85bdb2147ba131d", } - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=AsyncMock() ) @@ -582,7 +582,7 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi # Intentionally omit model_name to mimic Vertex issue. } - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=AsyncMock() ) @@ -658,7 +658,7 @@ async def test_error_file_id_for_failed_batch(): "unified_batch_id": "litellm_proxy;model_id:test-model-id;llm_batch_id:batch_abc123", } - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=AsyncMock() ) @@ -728,7 +728,7 @@ async def test_async_post_call_success_hook_twice_assert_no_unique_violation(): "unified_batch_id": "litellm_proxy;model_id:12345679;llm_batch_id:batch_685c5e5d63988190b85bdb2147ba131d", } - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -1104,7 +1104,7 @@ def test_get_file_ids_from_responses_tools(): Test that get_file_ids_from_responses_tools correctly extracts file IDs from the tools parameter. """ - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) @@ -1127,7 +1127,7 @@ def test_get_file_ids_from_responses_tools_multiple_tools(): """ Test that get_file_ids_from_responses_tools handles multiple tools. """ - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) @@ -1161,7 +1161,7 @@ def test_get_file_ids_from_responses_tools_empty(): """ Test that get_file_ids_from_responses_tools handles empty or None tools. """ - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) @@ -1194,7 +1194,7 @@ async def test_check_file_ids_access_with_unified_file_ids(): prisma_client = AsyncMock() internal_usage_cache = MagicMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1231,7 +1231,7 @@ async def test_check_file_ids_access_denied(): prisma_client = AsyncMock() internal_usage_cache = MagicMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1268,7 +1268,7 @@ async def test_check_file_ids_access_with_regular_files_only(): prisma_client = AsyncMock() internal_usage_cache = MagicMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1306,7 +1306,7 @@ async def test_completion_with_file_access_check(): internal_usage_cache = MagicMock() internal_usage_cache.async_get_cache = AsyncMock(return_value=None) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1366,7 +1366,7 @@ async def test_responses_with_file_access_check(): internal_usage_cache = MagicMock() internal_usage_cache.async_get_cache = AsyncMock(return_value=None) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1431,7 +1431,7 @@ async def test_store_unified_file_id_with_none_file_object(): internal_usage_cache = MagicMock() internal_usage_cache.async_set_cache = AsyncMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1466,7 +1466,7 @@ async def test_store_unified_file_id_updates_file_metadata_on_existing_row(): internal_usage_cache = MagicMock() internal_usage_cache.async_set_cache = AsyncMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1545,7 +1545,7 @@ async def test_afile_delete_returns_provider_response_when_stored_file_object_no ) internal_usage_cache.async_set_cache = AsyncMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1588,7 +1588,7 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none(): prisma_client = AsyncMock() internal_usage_cache = MagicMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1642,7 +1642,7 @@ async def test_afile_retrieve_raises_error_when_no_router_and_file_object_none() prisma_client = AsyncMock() internal_usage_cache = MagicMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1676,7 +1676,7 @@ async def test_afile_retrieve_returns_stored_file_object_when_exists(): prisma_client = AsyncMock() internal_usage_cache = MagicMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1713,7 +1713,7 @@ async def test_afile_retrieve_raises_error_for_non_managed_file(): prisma_client = AsyncMock() internal_usage_cache = MagicMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=internal_usage_cache, prisma_client=prisma_client, ) @@ -1773,7 +1773,7 @@ async def test_list_batches_from_managed_objects_table(): batch_record_2, ] - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -1804,7 +1804,7 @@ async def test_list_batches_from_managed_objects_table_empty_list(): prisma_client = AsyncMock() prisma_client.db.litellm_managedobjecttable.find_many.return_value = [] - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -1897,7 +1897,7 @@ async def test_list_batches_registers_and_returns_unified_output_file_ids(): ) prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -1977,7 +1977,7 @@ async def test_list_batches_resolves_existing_managed_rows_without_minting(): ) prisma_client.db.litellm_managedfiletable.find_first = AsyncMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2002,7 +2002,7 @@ async def test_list_batches_caps_page_size_at_100(): prisma_client = AsyncMock() prisma_client.db.litellm_managedobjecttable.find_many.return_value = [] - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2024,7 +2024,7 @@ async def test_list_batches_from_managed_objects_table_provider_filter_raises_ex prisma_client = AsyncMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2050,7 +2050,7 @@ async def test_list_batches_from_managed_objects_table_target_model_name_filter_ prisma_client = AsyncMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2108,7 +2108,7 @@ async def test_list_batches_from_managed_objects_table_filters_by_created_by(): } ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2158,7 +2158,7 @@ async def test_list_batches_pagination_uses_unified_object_id_cursor(): prisma_client.db.litellm_managedobjecttable.find_first.return_value = MagicMock() prisma_client.db.litellm_managedobjecttable.find_many.return_value = [] - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2255,7 +2255,7 @@ async def test_list_batches_pagination_walks_all_pages_without_loops_or_gaps(): side_effect=fake_find_first ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) user = UserAPIKeyAuth(user_id="test-user") @@ -2376,7 +2376,7 @@ async def test_list_batches_pagination_stable_when_created_at_ties(): side_effect=fake_find_first ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) user = UserAPIKeyAuth(user_id="test-user") @@ -2491,7 +2491,7 @@ async def test_list_batches_rejects_unknown_after_cursor(): ) prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[]) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2526,7 +2526,7 @@ async def test_list_batches_treats_empty_after_as_no_cursor(): rows = [_managed_batch_row(i) for i in range(2)] prisma_client = _fake_managed_object_table(rows) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2568,7 +2568,7 @@ async def test_list_batches_rejects_after_cursor_owned_by_another_user(): ) prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[]) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2596,7 +2596,7 @@ async def test_list_batches_has_more_false_on_exactly_full_final_page(): rows = [_managed_batch_row(i) for i in range(4)] prisma_client = _fake_managed_object_table(rows) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2624,7 +2624,7 @@ async def test_list_batches_unparseable_row_does_not_truncate_pagination(): rows[2].file_object = "{ not valid json" prisma_client = _fake_managed_object_table(rows) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2653,7 +2653,7 @@ async def test_list_batches_fills_a_page_past_a_full_page_of_unparseable_rows(): corrupt_row.file_object = "{ not valid json" prisma_client = _fake_managed_object_table(rows) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2684,7 +2684,7 @@ async def test_list_batches_bounds_the_queries_a_deep_unparseable_run_costs(): ] prisma_client = _fake_managed_object_table(rows) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2708,7 +2708,7 @@ async def test_list_batches_reads_one_chunk_when_the_first_one_fills_the_page(): rows = [_managed_batch_row(index) for index in range(_DEEP_BATCH_SCAN_ROW_COUNT)] prisma_client = _fake_managed_object_table(rows) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2748,7 +2748,7 @@ async def test_return_unified_file_id_includes_expires_at(): internal_usage_cache = MagicMock() - result = await _PROXY_LiteLLMManagedFiles.return_unified_file_id( + result = await PROXY_LiteLLMManagedFiles.return_unified_file_id( file_objects=[file_object], create_file_request=create_file_request, internal_usage_cache=internal_usage_cache, @@ -2789,7 +2789,7 @@ async def test_user_b_cannot_retrieve_user_a_batch(): batch_record.created_by = "user_a_id" prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2826,7 +2826,7 @@ async def test_user_b_cannot_cancel_user_a_batch(): batch_record.created_by = "user_a_id" prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2866,7 +2866,7 @@ async def test_user_a_can_retrieve_own_batch(): batch_record.created_by = "user_a_id" prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -2904,7 +2904,7 @@ async def test_user_b_cannot_retrieve_user_a_file(): file_record.created_by = "user_a_id" prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) @@ -2941,7 +2941,7 @@ async def test_user_b_cannot_download_user_a_file_content(): file_record.created_by = "user_a_id" prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) @@ -2978,7 +2978,7 @@ async def test_user_b_cannot_delete_user_a_file(): file_record.created_by = "user_a_id" prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) @@ -3028,7 +3028,7 @@ async def test_user_a_can_retrieve_own_file(): ) prisma_client.db.litellm_managedfiletable.find_first.return_value = file_record - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) @@ -3081,7 +3081,7 @@ async def test_list_batches_only_returns_user_own_batches(): # Mock database to only return User A's batches prisma_client.db.litellm_managedobjecttable.find_many.return_value = [batch_user_a] - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -3120,7 +3120,7 @@ async def test_same_user_different_keys_can_access_batch(): batch_record.created_by = "user_a_id" prisma_client.db.litellm_managedobjecttable.find_first.return_value = batch_record - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -3212,7 +3212,7 @@ async def test_team_b_cannot_access_team_a_provider_format_batch( prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -3253,7 +3253,7 @@ async def test_authorized_callers_can_access_provider_format_batch(caller_kwargs prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -3280,7 +3280,7 @@ async def test_provider_format_batch_without_ownership_row_stays_accessible(): prisma_client = AsyncMock() prisma_client.db.litellm_managedobjecttable.find_first.return_value = None - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -3308,7 +3308,7 @@ async def test_fine_tuning_provider_format_id_not_enforced(): prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -3345,7 +3345,7 @@ async def test_team_b_cannot_access_team_a_provider_format_file( prisma_client.db.litellm_managedfiletable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) @@ -3373,7 +3373,7 @@ async def test_same_team_can_access_provider_format_file(): prisma_client.db.litellm_managedfiletable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) @@ -3395,7 +3395,7 @@ async def test_provider_format_file_without_ownership_row_stays_accessible(): prisma_client = AsyncMock() prisma_client.db.litellm_managedfiletable.find_first.return_value = None - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client ) @@ -3417,7 +3417,7 @@ async def test_post_call_batch_create_stores_ownership_row(batch_id): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client ) @@ -3452,7 +3452,7 @@ async def test_post_call_batch_sync_does_not_claim_ownership(): prisma_client = AsyncMock() prisma_client.db.litellm_managedobjecttable.update_many.return_value = 0 - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client ) @@ -3477,7 +3477,7 @@ async def test_post_call_batch_sync_updates_existing_row(): prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client ) @@ -3511,7 +3511,7 @@ async def test_post_call_batch_sync_stores_output_file_ownership_from_batch_row( prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client ) @@ -3542,7 +3542,7 @@ async def test_post_call_batch_create_does_not_store_output_file_ownership(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client ) @@ -3592,7 +3592,7 @@ async def test_file_list_cursors_are_scoped_to_the_caller(): prisma_client = AsyncMock() prisma_client.db.litellm_managedfiletable.find_many.return_value = [] - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -3649,7 +3649,7 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): } prisma_client = AsyncMock() prisma_client.db.litellm_managedfiletable.find_many.return_value = [managed_row] - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client ) @@ -3671,7 +3671,7 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): async def test_list_user_batches_provider_filter_rejected_with_400(): from litellm.proxy._types import ProxyException, UserAPIKeyAuth - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) @@ -3691,7 +3691,7 @@ async def test_list_user_batches_provider_filter_rejected_with_400(): async def test_list_user_batches_target_model_names_filter_rejected_with_400(): from litellm.proxy._types import ProxyException, UserAPIKeyAuth - proxy_managed_files = _PROXY_LiteLLMManagedFiles( + proxy_managed_files = PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=MagicMock() ) diff --git a/tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py b/tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py index 7040aef73e5..aca98929584 100644 --- a/tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py +++ b/tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py @@ -16,10 +16,10 @@ from litellm.types.llms.openai import OpenAIFileObject def _make_managed_files_instance(): from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) - instance = _PROXY_LiteLLMManagedFiles( + instance = PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=MagicMock(), ) diff --git a/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py b/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py index 55a4d05946e..c48af1f6177 100644 --- a/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py +++ b/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py @@ -182,14 +182,14 @@ async def test_ensure_batch_response_uses_batch_owner_when_db_batch_object_prese @pytest.mark.asyncio async def test_registered_output_file_row_denies_cross_user_access(): from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) raw_output_file_id = "file-raw-output" prisma = MagicMock() prisma.db.litellm_managedfiletable.upsert = AsyncMock() prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) - managed_files = _PROXY_LiteLLMManagedFiles( + managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=prisma, ) diff --git a/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py b/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py index d3e668b8987..2049ff58f95 100644 --- a/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py +++ b/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py @@ -359,8 +359,8 @@ async def test_ensure_batch_response_returns_early_without_auth(): def _in_memory_managed_files(): - """Build a real _PROXY_LiteLLMManagedFiles whose prisma upsert hits an in-memory row.""" - from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + """Build a real PROXY_LiteLLMManagedFiles whose prisma upsert hits an in-memory row.""" + from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles store: dict = {} @@ -381,7 +381,7 @@ def _in_memory_managed_files(): cache.async_set_cache = AsyncMock() return ( - _PROXY_LiteLLMManagedFiles(internal_usage_cache=cache, prisma_client=prisma), + PROXY_LiteLLMManagedFiles(internal_usage_cache=cache, prisma_client=prisma), store, ) diff --git a/tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py b/tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py index 7ad564dc8f9..4e80e4fd6d7 100644 --- a/tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py +++ b/tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py @@ -30,15 +30,15 @@ def _make_unified_file_id() -> str: def _make_managed_files_with_no_db_record(): - """Create a _PROXY_LiteLLMManagedFiles where the DB returns None (file was deleted).""" + """Create a PROXY_LiteLLMManagedFiles where the DB returns None (file was deleted).""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) mock_prisma = MagicMock() mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) - return _PROXY_LiteLLMManagedFiles( + return PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=mock_prisma, ) @@ -64,7 +64,7 @@ async def test_should_raise_404_for_deleted_file(): async def test_should_allow_owner_access_when_record_exists(): """Baseline: file owner can access their own file.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) unified_file_id = _make_unified_file_id() @@ -77,7 +77,7 @@ async def test_should_allow_owner_access_when_record_exists(): return_value=mock_db_record ) - managed_files = _PROXY_LiteLLMManagedFiles( + managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=mock_prisma, ) @@ -93,7 +93,7 @@ async def test_should_allow_owner_access_when_record_exists(): async def test_should_block_different_user_when_record_exists(): """Baseline: different user cannot access another user's file.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) unified_file_id = _make_unified_file_id() @@ -106,7 +106,7 @@ async def test_should_block_different_user_when_record_exists(): return_value=mock_db_record ) - managed_files = _PROXY_LiteLLMManagedFiles( + managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=mock_prisma, ) diff --git a/tests/unit/enterprise/proxy/test_file_deletion_blocking.py b/tests/unit/enterprise/proxy/test_file_deletion_blocking.py index 852077dcf0c..b7acdad35a1 100644 --- a/tests/unit/enterprise/proxy/test_file_deletion_blocking.py +++ b/tests/unit/enterprise/proxy/test_file_deletion_blocking.py @@ -60,7 +60,7 @@ def _make_managed_files_instance_with_batches( file_created_by: str = "user-A", ): """ - Create a _PROXY_LiteLLMManagedFiles instance with mocked DB and batches. + Create a PROXY_LiteLLMManagedFiles instance with mocked DB and batches. Args: file_id: The unified file ID @@ -68,7 +68,7 @@ def _make_managed_files_instance_with_batches( file_created_by: The user who created the file """ from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) # Mock file record @@ -102,7 +102,7 @@ def _make_managed_files_instance_with_batches( }) mock_cache.async_set_cache = AsyncMock() - instance = _PROXY_LiteLLMManagedFiles( + instance = PROXY_LiteLLMManagedFiles( internal_usage_cache=mock_cache, prisma_client=mock_prisma, ) @@ -115,10 +115,10 @@ def _make_managed_files_instance_with_batches( def test_is_batch_polling_enabled_when_job_registered(): """Test that batch polling is detected as enabled when scheduler job is registered.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) - instance = _PROXY_LiteLLMManagedFiles( + instance = PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=MagicMock(), ) @@ -135,10 +135,10 @@ def test_is_batch_polling_enabled_when_job_registered(): def test_is_batch_polling_disabled_when_job_not_registered(): """Test that batch polling is detected as disabled when scheduler job is not registered.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) - instance = _PROXY_LiteLLMManagedFiles( + instance = PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=MagicMock(), ) @@ -154,10 +154,10 @@ def test_is_batch_polling_disabled_when_job_not_registered(): def test_is_batch_polling_disabled_when_no_scheduler(): """Test that batch polling is detected as disabled when scheduler is not available.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) - instance = _PROXY_LiteLLMManagedFiles( + instance = PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=MagicMock(), ) diff --git a/tests/unit/enterprise/proxy/test_managed_files_access_check.py b/tests/unit/enterprise/proxy/test_managed_files_access_check.py index ad46798b788..62302446a29 100644 --- a/tests/unit/enterprise/proxy/test_managed_files_access_check.py +++ b/tests/unit/enterprise/proxy/test_managed_files_access_check.py @@ -39,9 +39,9 @@ def _make_managed_files_instance( unified_file_id: str, file_team_id=None, ): - """Create a _PROXY_LiteLLMManagedFiles with a mocked DB that returns a file owned by file_created_by.""" + """Create a PROXY_LiteLLMManagedFiles with a mocked DB that returns a file owned by file_created_by.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) mock_db_record = MagicMock() @@ -53,7 +53,7 @@ def _make_managed_files_instance( return_value=mock_db_record ) - instance = _PROXY_LiteLLMManagedFiles( + instance = PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=mock_prisma, ) @@ -176,7 +176,7 @@ def _make_managed_files_instance_with_object_store(): """Managed-files hook backed by an in-memory stand-in for the managed object table, so create and retrieve exercise the same stored row.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) store = {} @@ -194,7 +194,7 @@ def _make_managed_files_instance_with_object_store(): ) return ( - _PROXY_LiteLLMManagedFiles( + PROXY_LiteLLMManagedFiles( internal_usage_cache=DualCache(), prisma_client=mock_prisma, ), diff --git a/tests/unit/enterprise/proxy/test_managed_files_hook.py b/tests/unit/enterprise/proxy/test_managed_files_hook.py index d99ea5ab445..3ea3b97e1fe 100644 --- a/tests/unit/enterprise/proxy/test_managed_files_hook.py +++ b/tests/unit/enterprise/proxy/test_managed_files_hook.py @@ -177,15 +177,15 @@ def _make_managed_files_over_rows(rows): def _make_managed_files_instance(): - """Create a _PROXY_LiteLLMManagedFiles with storage methods mocked out.""" + """Create a PROXY_LiteLLMManagedFiles with storage methods mocked out.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) mock_cache = MagicMock() mock_prisma = MagicMock() - instance = _PROXY_LiteLLMManagedFiles( + instance = PROXY_LiteLLMManagedFiles( internal_usage_cache=mock_cache, prisma_client=mock_prisma, ) @@ -1193,7 +1193,7 @@ def _managed_deletion_file_id(provider_file_id): def _managed_files_with_deletion_row(unified_file_id, provider_file_id, file_object): from litellm.caching import DualCache from litellm.models.managed_files import LiteLLM_ManagedFileTable - from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles row = LiteLLM_ManagedFileTable( unified_file_id=unified_file_id, @@ -1205,7 +1205,7 @@ def _managed_files_with_deletion_row(unified_file_id, provider_file_id, file_obj find_first=AsyncMock(return_value=row), delete=AsyncMock(), ) - return _PROXY_LiteLLMManagedFiles( + return PROXY_LiteLLMManagedFiles( internal_usage_cache=DualCache(), prisma_client=MagicMock(db=MagicMock(litellm_managedfiletable=table)), ), table @@ -1377,10 +1377,10 @@ async def test_afile_content_error_reports_unified_id_not_provider_uri(): def _make_real_managed_files_instance(): - """Create a _PROXY_LiteLLMManagedFiles with a real store_unified_file_id but + """Create a PROXY_LiteLLMManagedFiles with a real store_unified_file_id but an AsyncMock prisma client, so the DB write path itself can be asserted.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) mock_cache = MagicMock() @@ -1395,7 +1395,7 @@ def _make_real_managed_files_instance(): ) return ( - _PROXY_LiteLLMManagedFiles( + PROXY_LiteLLMManagedFiles( internal_usage_cache=mock_cache, prisma_client=mock_prisma, ), @@ -1407,7 +1407,7 @@ def _make_object_store_instance(): """A real store_unified_object_id over an AsyncMock prisma client, so both the upsert and the update-only write path can be asserted.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) mock_cache = MagicMock() @@ -1418,7 +1418,7 @@ def _make_object_store_instance(): mock_prisma.db.litellm_managedobjecttable.update_many = AsyncMock() return ( - _PROXY_LiteLLMManagedFiles( + PROXY_LiteLLMManagedFiles( internal_usage_cache=mock_cache, prisma_client=mock_prisma, ), @@ -1955,7 +1955,7 @@ async def test_afile_delete_bedrock_unified_id_end_to_end(monkeypatch): @pytest.mark.asyncio async def test_afile_delete_storage_backed_row_deletes_stored_content_not_provider_files(): - from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles from openai.types import FileDeleted from litellm.caching import DualCache @@ -1973,7 +1973,7 @@ async def test_afile_delete_storage_backed_row_deletes_stored_content_not_provid ) file_table = MagicMock(find_first=AsyncMock(return_value=row), delete=AsyncMock()) content_table = MagicMock(delete=AsyncMock()) - managed_files = _PROXY_LiteLLMManagedFiles( + managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=DualCache(), prisma_client=MagicMock( db=MagicMock(litellm_managedfiletable=file_table, litellm_managedfilecontenttable=content_table) @@ -1998,7 +1998,7 @@ async def test_afile_delete_storage_backed_row_deletes_stored_content_not_provid @pytest.mark.asyncio async def test_afile_content_storage_backed_row_returns_stored_bytes_not_provider_content(): - from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles from prisma import Base64 from litellm.caching import DualCache @@ -2017,7 +2017,7 @@ async def test_afile_content_storage_backed_row_returns_stored_bytes_not_provide ) file_table = MagicMock(find_first=AsyncMock(return_value=row)) content_table = MagicMock(find_unique=AsyncMock(return_value=MagicMock(content=Base64.encode(stored_bytes)))) - managed_files = _PROXY_LiteLLMManagedFiles( + managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=DualCache(), prisma_client=MagicMock( db=MagicMock(litellm_managedfiletable=file_table, litellm_managedfilecontenttable=content_table) @@ -2070,12 +2070,12 @@ async def test_store_unified_object_id_batch_processed_is_written_only_when_aske @pytest.mark.asyncio async def test_store_unified_file_id_caches_the_storage_location_the_db_row_gets(): - from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles from litellm.caching import DualCache file_table = MagicMock(upsert=AsyncMock(), find_first=AsyncMock(side_effect=AssertionError("cache miss"))) - managed_files = _PROXY_LiteLLMManagedFiles( + managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=DualCache(), prisma_client=MagicMock(db=MagicMock(litellm_managedfiletable=file_table)), ) diff --git a/tests/unit/google_genai/test_google_genai_adapter.py b/tests/unit/google_genai/test_google_genai_adapter.py index 761ab7bac89..38292f471c8 100644 --- a/tests/unit/google_genai/test_google_genai_adapter.py +++ b/tests/unit/google_genai/test_google_genai_adapter.py @@ -847,7 +847,7 @@ def test_api_base_and_api_key_passthrough(function_name, is_async, is_stream): import asyncio import unittest.mock - litellm._turn_on_debug() + litellm.turn_on_debug() # Import the specific function being tested if function_name == "generate_content": diff --git a/tests/unit/images/test_main.py b/tests/unit/images/test_main.py index 29be4bd7a52..e78323868f0 100644 --- a/tests/unit/images/test_main.py +++ b/tests/unit/images/test_main.py @@ -72,7 +72,7 @@ def test_image_edit_prices_a_vertex_deployment_at_its_configured_location( vertex_location=location, litellm_logging_obj=logging_obj, ) - return logging_obj._response_cost_calculator(result=response) + return logging_obj.response_cost_calculator(result=response) assert cost_at("global") == pytest.approx(0.04) assert cost_at("us-central1") == pytest.approx(0.044) diff --git a/tests/unit/integrations/SlackAlerting/test_slack_alerting_utils.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting_utils.py index b02dbea64b8..04c2dc7786f 100644 --- a/tests/unit/integrations/SlackAlerting/test_slack_alerting_utils.py +++ b/tests/unit/integrations/SlackAlerting/test_slack_alerting_utils.py @@ -4,7 +4,7 @@ import pytest # Adds the grandparent directory to sys.path to allow importing project modules import litellm -from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert +from litellm.integrations.SlackAlerting.utils import add_langfuse_trace_id_to_alert from litellm.litellm_core_utils.logging_callback_manager import LoggingCallbackManager @@ -21,7 +21,7 @@ async def test_langfuse_not_initialized_returns_none_early(): request_data = {"litellm_logging_obj": MagicMock(), "trace_id": "test-trace-id"} # Call the function - result = await _add_langfuse_trace_id_to_alert(request_data) + result = await add_langfuse_trace_id_to_alert(request_data) # Should return None early without processing request_data assert result is None @@ -40,10 +40,10 @@ async def test_langfuse_trace_url_uses_the_request_host_without_building_a_logge monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) logging_obj = MagicMock() - logging_obj._get_trace_id.return_value = "abc123" + logging_obj.get_trace_id.return_value = "abc123" logging_obj.standard_callback_dynamic_params = {"langfuse_host": "http://127.0.0.1:1"} - result = await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) + result = await add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) assert result == "http://127.0.0.1:1/trace/abc123" assert litellm.initialized_langfuse_clients == 0 @@ -54,10 +54,10 @@ async def test_langfuse_trace_url_falls_back_to_the_env_host(monkeypatch): monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) monkeypatch.setenv("LANGFUSE_HOST", "langfuse.internal:3000") logging_obj = MagicMock() - logging_obj._get_trace_id.return_value = "abc123" + logging_obj.get_trace_id.return_value = "abc123" logging_obj.standard_callback_dynamic_params = {} - assert await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) == ( + assert await add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) == ( "http://langfuse.internal:3000/trace/abc123" ) @@ -78,10 +78,10 @@ async def test_langfuse_trace_url_when_callback_registered_as_logger_instance(mo monkeypatch.setattr(litellm, "callbacks", []) monkeypatch.setenv("LANGFUSE_HOST", "http://env-host.invalid") logging_obj = MagicMock() - logging_obj._get_trace_id.return_value = "trace-from-instance" + logging_obj.get_trace_id.return_value = "trace-from-instance" logging_obj.standard_callback_dynamic_params = {} - result = await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) + result = await add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) assert result == "http://127.0.0.1:1/trace/trace-from-instance" @@ -103,10 +103,10 @@ async def test_langfuse_trace_url_when_prompt_management_is_the_registered_callb monkeypatch.setattr(litellm, "callbacks", [prompt_callback]) monkeypatch.setenv("LANGFUSE_HOST", "http://env-host.invalid") logging_obj = MagicMock() - logging_obj._get_trace_id.return_value = "trace-from-prompt-callback" + logging_obj.get_trace_id.return_value = "trace-from-prompt-callback" logging_obj.standard_callback_dynamic_params = {} - result = await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) + result = await add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) assert result == "http://127.0.0.1:2/trace/trace-from-prompt-callback" @@ -116,7 +116,7 @@ async def test_langfuse_trace_url_absent_when_trace_id_never_arrives(monkeypatch monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) monkeypatch.setattr("litellm.integrations.SlackAlerting.utils.asyncio.sleep", AsyncMock()) logging_obj = MagicMock() - logging_obj._get_trace_id.return_value = None + logging_obj.get_trace_id.return_value = None logging_obj.standard_callback_dynamic_params = {"langfuse_host": "http://127.0.0.1:1"} - assert await _add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) is None + assert await add_langfuse_trace_id_to_alert({"litellm_logging_obj": logging_obj}) is None diff --git a/tests/unit/integrations/datadog/test_datadog_team_handler.py b/tests/unit/integrations/datadog/test_datadog_team_handler.py index 09d6f51e0a8..925a549b19e 100644 --- a/tests/unit/integrations/datadog/test_datadog_team_handler.py +++ b/tests/unit/integrations/datadog/test_datadog_team_handler.py @@ -221,7 +221,7 @@ class TestDataDogHandler: assert result.DD_API_KEY == "team_key" assert "eu1.datadoghq.com" in result.intake_url - def test_request_blocked_callback_params_includes_dd(self): + def test_team_callback_params_are_blocked_for_requests(self): """DD params should be blocked from request-level metadata (security).""" from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( _request_blocked_callback_params, diff --git a/tests/unit/integrations/focus/test_export_engine.py b/tests/unit/integrations/focus/test_export_engine.py new file mode 100644 index 00000000000..ab40ab8196c --- /dev/null +++ b/tests/unit/integrations/focus/test_export_engine.py @@ -0,0 +1,70 @@ +from datetime import datetime, timezone + +import polars as pl + +from litellm.integrations.focus.database import FocusLiteLLMDatabase +from litellm.integrations.focus.destinations import FocusTimeWindow +from litellm.integrations.focus.destinations.factory import FocusDestinationFactory +from litellm.integrations.focus.export_engine import FocusExportEngine +from litellm.integrations.focus.serializers import FocusParquetSerializer +from litellm.integrations.focus.transformer import FocusTransformer + + +def _vantage_csv_engine() -> FocusExportEngine: + return FocusExportEngine( + provider="vantage", + export_format="csv", + prefix="exports", + destination_config={"api_key": "test-key", "integration_token": "test-token"}, + ) + + +def test_legacy_engine_properties_forward_to_public_storage(): + engine = _vantage_csv_engine() + + replacement_destination = FocusDestinationFactory.create( + provider="vantage", + prefix="replacement", + config={"api_key": "test-key", "integration_token": "test-token"}, + ) + engine._destination = replacement_destination + assert engine._destination is engine.destination is replacement_destination + + replacement_serializer = FocusParquetSerializer() + engine._serializer = replacement_serializer + assert engine._serializer is engine.serializer is replacement_serializer + + replacement_transformer = FocusTransformer() + engine._transformer = replacement_transformer + assert engine._transformer is engine.transformer is replacement_transformer + + replacement_database = FocusLiteLLMDatabase() + engine._database = replacement_database + assert engine._database is engine.database is replacement_database + + +def test_build_filename_uses_the_window_and_serializer_format(): + engine = _vantage_csv_engine() + window = FocusTimeWindow( + start_time=datetime(2024, 1, 2, 3, 4, 5, tzinfo=timezone.utc), + end_time=datetime(2024, 1, 2, 4, 5, 6, tzinfo=timezone.utc), + frequency="hourly", + ) + + assert engine.build_filename(window) == "usage_20240102T030405Z_20240102T040506Z.csv" + + +def test_export_aggregates_compute_values_and_handle_missing_columns(): + frame = pl.DataFrame( + { + "spend": [1.5, 2.5, 1.0], + "team_id": ["team-a", "team-a", "team-b"], + } + ) + + assert FocusExportEngine.sum_column(frame, "spend") == 5.0 + assert FocusExportEngine.count_unique(frame, "team_id") == 2 + assert FocusExportEngine.sum_column(frame, "missing") == 0.0 + assert FocusExportEngine.count_unique(frame, "missing") == 0 + assert FocusExportEngine.sum_column(pl.DataFrame(), "spend") == 0.0 + assert FocusExportEngine.count_unique(pl.DataFrame(), "team_id") == 0 diff --git a/tests/unit/integrations/focus/test_mavvrik_destination.py b/tests/unit/integrations/focus/test_mavvrik_destination.py index 574c6186e07..1be72b4409b 100644 --- a/tests/unit/integrations/focus/test_mavvrik_destination.py +++ b/tests/unit/integrations/focus/test_mavvrik_destination.py @@ -482,8 +482,8 @@ async def test_export_window_passes_max_rows_as_limit(monkeypatch): db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) # empty deliver engine_mock = MagicMock() - engine_mock._database = db_mock - engine_mock._destination.deliver = AsyncMock() + engine_mock.database = db_mock + engine_mock.destination.deliver = AsyncMock() logger._engine = engine_mock window = FocusTimeWindow( @@ -530,8 +530,8 @@ async def test_run_scheduled_export_catches_up_missed_dates(): db_mock = MagicMock() db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) engine_mock = MagicMock() - engine_mock._database = db_mock - engine_mock._destination = dest_mock + engine_mock.database = db_mock + engine_mock.destination = dest_mock logger._engine = engine_mock await logger._run_scheduled_export() @@ -571,8 +571,8 @@ async def test_run_scheduled_export_no_catchup_when_marker_is_current(): db_mock = MagicMock() db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) engine_mock = MagicMock() - engine_mock._database = db_mock - engine_mock._destination = dest_mock + engine_mock.database = db_mock + engine_mock.destination = dest_mock logger._engine = engine_mock await logger._run_scheduled_export() @@ -602,8 +602,8 @@ async def test_run_scheduled_export_skips_catchup_when_marker_is_unparseable(): db_mock = MagicMock() db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) engine_mock = MagicMock() - engine_mock._database = db_mock - engine_mock._destination = dest_mock + engine_mock.database = db_mock + engine_mock.destination = dest_mock logger._engine = engine_mock await logger._run_scheduled_export() @@ -729,8 +729,8 @@ async def test_catchup_capped_at_max_catchup_days(): db_mock = MagicMock() db_mock.get_usage_data = AsyncMock(return_value=pl.DataFrame()) engine_mock = MagicMock() - engine_mock._database = db_mock - engine_mock._destination = dest_mock + engine_mock.database = db_mock + engine_mock.destination = dest_mock logger._engine = engine_mock await logger._run_scheduled_export() diff --git a/tests/unit/integrations/gitlab/test_gitlab_prompt_manager.py b/tests/unit/integrations/gitlab/test_gitlab_prompt_manager.py index 120cc877b51..4240617a7d0 100644 --- a/tests/unit/integrations/gitlab/test_gitlab_prompt_manager.py +++ b/tests/unit/integrations/gitlab/test_gitlab_prompt_manager.py @@ -568,7 +568,7 @@ def test_list_templates_returns_encoded_ids(manager): def test_load_prompt_from_gitlab_parses_metadata(manager, mock_gitlab_client): - manager._load_prompt_from_gitlab("gitlab::hello") + manager.load_prompt_from_gitlab("gitlab::hello") assert "gitlab::hello" in manager.prompts tmpl = manager.prompts["gitlab::hello"] @@ -578,7 +578,7 @@ def test_load_prompt_from_gitlab_parses_metadata(manager, mock_gitlab_client): def test_render_template_renders_jinja(manager, mock_gitlab_client): - manager._load_prompt_from_gitlab("gitlab::hello") + manager.load_prompt_from_gitlab("gitlab::hello") output = manager.render_template("gitlab::hello", {"name": "Prishu"}) assert "Hello Prishu" in output @@ -589,7 +589,7 @@ def test_get_template_returns_none_if_not_loaded(manager): def test_repo_path_conversion(manager): raw = "gitlab::nested::sub" - repo_path = manager._id_to_repo_path(raw) + repo_path = manager.id_to_repo_path(raw) assert repo_path.endswith("nested/sub.prompt") # Ensure decode/encode reversibility decoded = manager._repo_path_to_id(repo_path) @@ -689,7 +689,7 @@ class FakeTemplateManager: """ def __init__(self, prompts_path="prompts"): - # simulate a configured prompts folder (affects _id_to_repo_path) + # simulate a configured prompts folder (affects id_to_repo_path) self.prompts_path = prompts_path.strip("/") self.prompts = {} # id -> GitLabPromptTemplate @@ -700,7 +700,7 @@ class FakeTemplateManager: def list_templates(self, *, recursive: bool = True): return list(self._discoverable_ids) - def _load_prompt_from_gitlab(self, pid, ref=None): + def load_prompt_from_gitlab(self, pid, ref=None): # Pretend we fetched and parsed a file; add a basic template if not present if pid not in self.prompts: self.prompts[pid] = GitLabPromptTemplate( @@ -712,7 +712,7 @@ class FakeTemplateManager: def get_template(self, pid): return self.prompts.get(pid) - def _id_to_repo_path(self, pid): + def id_to_repo_path(self, pid): base = f"{self.prompts_path}/" if self.prompts_path else "" return f"{base}{pid}.prompt" @@ -758,8 +758,8 @@ def test_cache_load_all_encodes_ids_and_populates_maps(mock_pm_cls, fake_manager assert set(result.keys()) == {encode_prompt_id("a"), encode_prompt_id("sub/b")} # Files map built with full repo paths - expect_a_path = tm._id_to_repo_path("a") - expect_b_path = tm._id_to_repo_path("sub/b") + expect_a_path = tm.id_to_repo_path("a") + expect_b_path = tm.id_to_repo_path("sub/b") assert cache.list_files() == [expect_a_path, expect_b_path] # IDs list is the encoded IDs @@ -829,7 +829,7 @@ def test_cache_skips_when_template_missing_even_after_reload_attempt( # Always return None to trigger the continue path return None - def _load_prompt_from_gitlab(self, pid, ref=None): + def load_prompt_from_gitlab(self, pid, ref=None): # Pretend to load, but still don't populate prompts so get_template stays None pass @@ -855,8 +855,8 @@ def test_cache_get_by_file_returns_exact_entry(mock_pm_cls, fake_managers): cache = GitLabPromptCache({"project": "g/s/r", "access_token": "tkn"}) cache.load_all() - alpha_path = tm._id_to_repo_path("alpha") - beta_path = tm._id_to_repo_path("nested/beta") + alpha_path = tm.id_to_repo_path("alpha") + beta_path = tm.id_to_repo_path("nested/beta") alpha = cache.get_by_file(alpha_path) beta = cache.get_by_file(beta_path) diff --git a/tests/unit/integrations/mavvrik_focus/test_mavvrik_focus_logger.py b/tests/unit/integrations/mavvrik_focus/test_mavvrik_focus_logger.py index cd21807e887..1dcc4567b2c 100644 --- a/tests/unit/integrations/mavvrik_focus/test_mavvrik_focus_logger.py +++ b/tests/unit/integrations/mavvrik_focus/test_mavvrik_focus_logger.py @@ -49,11 +49,11 @@ async def test_export_window_delivers_empty_payload_for_empty_export( destination = MagicMock() destination.deliver = AsyncMock() engine = MagicMock() - engine._database = database - engine._transformer = transformer - engine._serializer = serializer - engine._destination = destination - engine._build_filename.return_value = "metrics.csv" + engine.database = database + engine.transformer = transformer + engine.serializer = serializer + engine.destination = destination + engine.build_filename.return_value = "metrics.csv" logger = MavvrikFocusLogger() logger._engine = engine window = _window() diff --git a/tests/unit/integrations/newrelic/test_newrelic_team_handler.py b/tests/unit/integrations/newrelic/test_newrelic_team_handler.py index f4460a615df..1043da4be51 100644 --- a/tests/unit/integrations/newrelic/test_newrelic_team_handler.py +++ b/tests/unit/integrations/newrelic/test_newrelic_team_handler.py @@ -139,7 +139,7 @@ class TestNewRelicHandler: assert result_us.metric_api_url == US_ENDPOINT assert result_eu.metric_api_url == EU_ENDPOINT - def test_request_blocked_callback_params_includes_newrelic(self): + def test_team_callback_params_are_blocked_for_requests(self): from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( _request_blocked_callback_params, ) diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index 4c9fd3cfdf2..70039646b4b 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -3705,9 +3705,9 @@ class TestPromptCacheBreakpointCapability: bundled = os.path.join(os.path.dirname(litellm.__file__), "model_prices_and_context_window_backup.json") with open(bundled) as handle: monkeypatch.setattr(litellm, "model_cost", json.load(handle)) - litellm.utils._cached_get_model_info_helper.cache_clear() + litellm.utils.cached_get_model_info_helper.cache_clear() yield - litellm.utils._cached_get_model_info_helper.cache_clear() + litellm.utils.cached_get_model_info_helper.cache_clear() def test_listed_model_uses_the_model_map_flag(self, monkeypatch): diff --git a/tests/unit/integrations/test_prometheus_labels.py b/tests/unit/integrations/test_prometheus_labels.py index d598d59e183..8a1f5f5a0a8 100644 --- a/tests/unit/integrations/test_prometheus_labels.py +++ b/tests/unit/integrations/test_prometheus_labels.py @@ -581,9 +581,9 @@ def test_prometheus_label_value_sanitization(): def test_prometheus_label_value_sanitization_unicode_paragraph_separator(): """Test that U+2029 (Paragraph Separator) is also stripped.""" - from litellm.types.integrations.prometheus import _sanitize_prometheus_label_value + from litellm.types.integrations.prometheus import sanitize_prometheus_label_value - result = _sanitize_prometheus_label_value("model\u2029name") + result = sanitize_prometheus_label_value("model\u2029name") assert result == "modelname" assert "\u2029" not in result @@ -592,20 +592,20 @@ def test_prometheus_label_value_sanitization_unicode_paragraph_separator(): def test_prometheus_label_value_sanitization_none(): """Test that None values pass through unchanged.""" - from litellm.types.integrations.prometheus import _sanitize_prometheus_label_value + from litellm.types.integrations.prometheus import sanitize_prometheus_label_value - assert _sanitize_prometheus_label_value(None) is None + assert sanitize_prometheus_label_value(None) is None print("✅ None values pass through unchanged") def test_prometheus_label_value_sanitization_non_string_types(): """Test that non-string values (int, bool, etc.) are coerced to str.""" - from litellm.types.integrations.prometheus import _sanitize_prometheus_label_value + from litellm.types.integrations.prometheus import sanitize_prometheus_label_value - assert _sanitize_prometheus_label_value(200) == "200" - assert _sanitize_prometheus_label_value(True) == "True" - assert _sanitize_prometheus_label_value(3.14) == "3.14" + assert sanitize_prometheus_label_value(200) == "200" + assert sanitize_prometheus_label_value(True) == "True" + assert sanitize_prometheus_label_value(3.14) == "3.14" print("✅ Non-string values are coerced to str") diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 088247c2ea4..ea628794c2c 100644 --- a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -2419,17 +2419,17 @@ def test_service_tier_suffixes_constant_in_sync_with_enum(): def test_get_cost_per_unit_falls_back_from_service_tier_key_to_base(): - from litellm.litellm_core_utils.llm_cost_calc.utils import _get_cost_per_unit + from litellm.litellm_core_utils.llm_cost_calc.utils import get_cost_per_unit model_info = {"input_cost_per_token": 2e-6} # service-tier key is absent -> falls back to the base key - assert _get_cost_per_unit(model_info, "input_cost_per_token_priority") == 2e-6 + assert get_cost_per_unit(model_info, "input_cost_per_token_priority") == 2e-6 # service-tier key present -> used directly, no fallback model_info_direct = { "input_cost_per_token_priority": 5e-6, "input_cost_per_token": 2e-6, } - assert _get_cost_per_unit(model_info_direct, "input_cost_per_token_priority") == 5e-6 + assert get_cost_per_unit(model_info_direct, "input_cost_per_token_priority") == 5e-6 def test_threshold_keys_exclude_service_tier_variants(): @@ -2860,7 +2860,7 @@ def test_token_type_cost_breakdown_openai_responses_api_cache_write_read( from litellm.responses.utils import ResponseAPILoggingUtils model = "gpt-5.6" - usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_usage) + usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(raw_usage) breakdown = get_token_type_cost_breakdown(model=model, custom_llm_provider="openai", usage=usage) diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_openai_cache_write_cost.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_openai_cache_write_cost.py index 18acfeda07d..6c333a83d14 100644 --- a/tests/unit/litellm_core_utils/llm_cost_calc/test_openai_cache_write_cost.py +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_openai_cache_write_cost.py @@ -64,7 +64,7 @@ def test_responses_api_cache_write_costs_the_same_as_chat(local_model_cost_map): cache_write_tokens = 12314 completion_tokens = 5 - responses_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( + responses_usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage( { "input_tokens": prompt_tokens, "output_tokens": completion_tokens, diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_responses_cache_cost_breakdown.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_responses_cache_cost_breakdown.py index 0d387233be7..b32f9f63073 100644 --- a/tests/unit/litellm_core_utils/llm_cost_calc/test_responses_cache_cost_breakdown.py +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_responses_cache_cost_breakdown.py @@ -39,7 +39,7 @@ def _logging_obj() -> Logging: def _responses_completion(cached_tokens: int, cache_write_tokens: int, fresh_tokens: int, output_tokens: int): input_tokens = cached_tokens + cache_write_tokens + fresh_tokens - usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( + usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage( { "input_tokens": input_tokens, "output_tokens": output_tokens, diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py index 41a2d19b8ab..15bea2a2772 100644 --- a/tests/unit/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py @@ -49,6 +49,27 @@ def test_web_search_cost_high(): ) +def test_get_web_search_options_preserves_explicit_none_error(): + with pytest.raises(TypeError): + StandardBuiltInToolCostTracking.get_web_search_options( + {"web_search_options": None, "tools": [{"type": "web_search_preview"}]} + ) + + +def test_get_web_search_options_accepts_non_list_tool_iterables(): + options = StandardBuiltInToolCostTracking.get_web_search_options( + {"tools": ({"type": "web_search_preview"},)} + ) + + assert options == {"type": "web_search_preview"} + + +def test_get_file_search_tool_call_does_not_validate_tool_payload(): + tool = {"type": "file_search", "vector_store_ids": "not-a-list"} + + assert StandardBuiltInToolCostTracking.get_file_search_tool_call({"tools": [tool]}) == tool + + # Test file search cost calculation def test_file_search_cost(): file_search = FileSearchTool(type="file_search") @@ -571,4 +592,3 @@ _BEDROCK_MANTLE_WEB_SEARCH_MODELS = ( _BEDROCK_MANTLE_WEB_SEARCH_RATE = 0.012 - diff --git a/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py b/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py index 8e46ae21de6..2865ffc13f1 100644 --- a/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py +++ b/tests/unit/litellm_core_utils/llm_response_utils/test_convert_dict_to_response.py @@ -4,8 +4,9 @@ import pytest from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( - _handle_invalid_parallel_tool_calls, - _should_convert_tool_call_to_json_mode, + handle_invalid_parallel_tool_calls, + should_convert_tool_call_to_json_mode, + safe_convert_created_field, convert_to_model_response_object, ) from litellm.types.utils import ( @@ -46,6 +47,20 @@ OPENAI_CUSTOM_TOOL_CALL_RESPONSE = { } +def test_safe_convert_created_field_preserves_large_integer_precision(): + created_value = 2**53 + 1 + + assert safe_convert_created_field(created_value) == created_value + + +def test_safe_convert_created_field_accepts_float_convertible_non_strings(): + class FloatConvertible: + def __float__(self) -> float: + return 1.5 + + assert safe_convert_created_field(FloatConvertible()) == 1 + + def test_convert_openai_custom_tool_call_response(): result = convert_to_model_response_object( response_object=OPENAI_CUSTOM_TOOL_CALL_RESPONSE, @@ -66,7 +81,7 @@ def test_should_convert_tool_call_to_json_mode_ignores_custom_tool_call(): custom={"name": "ApplyPatch", "input": "patch"}, ) assert ( - _should_convert_tool_call_to_json_mode( + should_convert_tool_call_to_json_mode( tool_calls=[custom_tool_call], convert_tool_call_to_json_mode=True, ) @@ -81,7 +96,7 @@ def test_should_convert_tool_call_to_json_mode_still_matches_response_format_too function=Function(name=RESPONSE_FORMAT_TOOL_NAME, arguments='{"answer": 4}'), ) assert ( - _should_convert_tool_call_to_json_mode( + should_convert_tool_call_to_json_mode( tool_calls=[response_format_call], convert_tool_call_to_json_mode=True, ) @@ -99,7 +114,7 @@ def test_handle_invalid_parallel_tool_calls_skips_custom_tool_calls(): type="function", function=Function(name="get_weather", arguments='{"city": "SF"}'), ) - result = _handle_invalid_parallel_tool_calls([custom_tool_call, function_tool_call]) + result = handle_invalid_parallel_tool_calls([custom_tool_call, function_tool_call]) assert result == [custom_tool_call, function_tool_call] diff --git a/tests/unit/litellm_core_utils/llm_response_utils/test_response_metadata.py b/tests/unit/litellm_core_utils/llm_response_utils/test_response_metadata.py index 554447f4273..00b0e2f6972 100644 --- a/tests/unit/litellm_core_utils/llm_response_utils/test_response_metadata.py +++ b/tests/unit/litellm_core_utils/llm_response_utils/test_response_metadata.py @@ -78,7 +78,7 @@ class TestCallbackDurationMs: """End-to-end: update_response_metadata should propagate callback_duration_ms.""" result = ModelResponse() logging_obj = self._make_logging_obj(callback_duration_ms=5.5, llm_api_duration_ms=800.0) - logging_obj._response_cost_calculator = MagicMock(return_value=0.001) + logging_obj.response_cost_calculator = MagicMock(return_value=0.001) logging_obj.litellm_call_id = "test-call-id" start = datetime.datetime(2025, 1, 1, 0, 0, 0) @@ -128,7 +128,7 @@ class TestDictResultsSkipMetadataUpdate: end_time=datetime.datetime(2025, 1, 1, 0, 0, 1), ) - logging_obj._response_cost_calculator.assert_not_called() + logging_obj.response_cost_calculator.assert_not_called() assert "_hidden_params" not in anthropic_response def test_update_response_metadata_keeps_timing_on_logging_obj_for_dict_results(self): @@ -152,7 +152,7 @@ class TestDictResultsSkipMetadataUpdate: logging_obj.set_response_timing_metrics.assert_called_once_with( {"_response_ms": 1000.0, "litellm_overhead_time_ms": 100.0} ) - logging_obj._response_cost_calculator.assert_not_called() + logging_obj.response_cost_calculator.assert_not_called() assert "_hidden_params" not in anthropic_response def test_update_response_metadata_keeps_timing_for_stream_wrapper_without_hidden_params(self): @@ -183,7 +183,7 @@ class TestDictResultsSkipMetadataUpdate: asyncio.run(drive()) logging_obj.set_response_timing_metrics.assert_called_once_with({"_response_ms": 250.0}) - logging_obj._response_cost_calculator.assert_not_called() + logging_obj.response_cost_calculator.assert_not_called() def test_update_response_metadata_leaves_logging_obj_alone_for_objects_with_hidden_params(self): """ModelResponse keeps carrying its own timing; the logging-object carrier is not written.""" @@ -191,7 +191,7 @@ class TestDictResultsSkipMetadataUpdate: logging_obj = MagicMock() logging_obj.model_call_details = {"llm_api_duration_ms": 900.0} logging_obj.caching_details = None - logging_obj._response_cost_calculator = MagicMock(return_value=0.001) + logging_obj.response_cost_calculator = MagicMock(return_value=0.001) logging_obj.litellm_call_id = "test-call-id" update_response_metadata( @@ -213,7 +213,7 @@ class TestDictResultsSkipMetadataUpdate: logging_obj = MagicMock() logging_obj.model_call_details = {"llm_api_duration_ms": 200.0} logging_obj.caching_details = None - logging_obj._response_cost_calculator = MagicMock(return_value=0.001) + logging_obj.response_cost_calculator = MagicMock(return_value=0.001) logging_obj.litellm_call_id = "test-call-id" update_response_metadata( diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 8d2e6b9fd0c..d91bc3b6201 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -22,15 +22,81 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( anthropic_messages_pt, convert_to_anthropic_tool_result, convert_to_gemini_tool_call_result, + encode_tool_call_id_with_signature, + function_call_prompt, + get_thought_signature_from_tool, get_tool_calls_from_response, make_valid_bedrock_tool_name, ollama_pt, + parse_mime_type, sanitize_messages_for_tool_calling, ) from litellm.types.llms.openai import ChatCompletionToolMessage from litellm.utils import validate_and_fix_openai_messages +def test_function_call_prompt_preserves_append_failure_for_non_string_content() -> None: + messages = [{"role": "system", "content": None}] + + with pytest.raises(AttributeError): + function_call_prompt(messages, []) + + +@pytest.mark.parametrize( + ("thought_signature", "expected"), + [ + ("encoded-signature", "call_123__thought__encoded-signature"), + (None, "call_123"), + ("", "call_123"), + ], +) +def test_encode_tool_call_id_with_signature(thought_signature, expected): + assert encode_tool_call_id_with_signature("call_123", thought_signature) == expected + + +@pytest.mark.parametrize( + ("tool", "expected"), + [ + ({"provider_specific_fields": {"thought_signature": "tool-signature"}}, "tool-signature"), + ( + {"function": {"provider_specific_fields": {"thought_signature": "function-signature"}}}, + "function-signature", + ), + ( + {"id": encode_tool_call_id_with_signature("call_123", "embedded-signature")}, + "embedded-signature", + ), + ({}, None), + ], +) +def test_get_thought_signature_from_tool(tool, expected): + assert get_thought_signature_from_tool(tool) == expected + + +@pytest.mark.parametrize( + ("base64_data", "expected"), + [ + ("data:image/png;base64,encoded-image", "image/png"), + ("data:application/pdf;base64,encoded-document", "application/pdf"), + ("not-a-data-url", None), + ], +) +def test_parse_mime_type(base64_data, expected): + assert parse_mime_type(base64_data) == expected + + +@pytest.mark.parametrize( + ("message", "expected_error"), + [ + ({"type": "file"}, "missing the required 'file' field"), + ({"type": "file", "file": {}}, "file_data and file_id cannot both be None"), + ], +) +def test_process_file_message_rejects_missing_file_data(message, expected_error): + with pytest.raises(litellm.BadRequestError, match=expected_error): + BedrockConverseMessagesProcessor.process_file_message(message) + + def _get_gemini_function_response_inline_data_parts(result): assert isinstance(result, list), "expected Gemini parts list" assert len(result) == 1, "multimodal function responses should stay in one part" @@ -324,7 +390,7 @@ def test_bedrock_validate_format_image_or_video(): # Test valid image formats valid_image_formats = ["png", "jpeg", "gif", "webp"] for format in valid_image_formats: - result = BedrockImageProcessor._validate_format(f"image/{format}", format) + result = BedrockImageProcessor.validate_format(f"image/{format}", format) assert result == format, f"Expected {format}, got {result}" # Test valid video formats @@ -340,7 +406,7 @@ def test_bedrock_validate_format_image_or_video(): "3gp", ] for format in valid_video_formats: - result = BedrockImageProcessor._validate_format(f"video/{format}", format) + result = BedrockImageProcessor.validate_format(f"video/{format}", format) assert result == format, f"Expected {format}, got {result}" # Test valid document formats @@ -352,7 +418,7 @@ def test_bedrock_validate_format_image_or_video(): } for mime, expected in valid_document_formats.items(): print("testing mime", mime, "expected", expected) - result = BedrockImageProcessor._validate_format(mime, mime.split("/")[1]) + result = BedrockImageProcessor.validate_format(mime, mime.split("/")[1]) assert result == expected, f"Expected {expected}, got {result}" diff --git a/tests/unit/litellm_core_utils/test_core_helpers.py b/tests/unit/litellm_core_utils/test_core_helpers.py index 9f48ea62a72..1f716b41acd 100644 --- a/tests/unit/litellm_core_utils/test_core_helpers.py +++ b/tests/unit/litellm_core_utils/test_core_helpers.py @@ -13,6 +13,7 @@ from litellm.litellm_core_utils.core_helpers import ( budget_reservation_from_metadata, drop_params_env_flag, drop_params_flag, + get_parent_otel_span_from_kwargs, get_or_create_metadata_bucket, get_provider_response_headers_from_hidden_params, map_finish_reason, @@ -27,6 +28,24 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.types.utils import ImageResponse, TranscriptionResponse +def test_parent_otel_span_reads_dict_like_metadata_by_index(): + from typing import cast + + span = object() + + class Metadata: + def __contains__(self, key: object) -> bool: + return key == "litellm_parent_otel_span" + + def __getitem__(self, key: str) -> object: + assert key == "litellm_parent_otel_span" + return span + + result = get_parent_otel_span_from_kwargs({"metadata": cast(object, Metadata())}) + + assert result is span + + @pytest.mark.parametrize("header", ("request-id", "x-request-id", "llm_provider-request-id")) def test_native_request_id_survives_stream_header_processing(header: str) -> None: processed: Final = process_response_headers(httpx.Headers({header: "req_native"})) diff --git a/tests/unit/litellm_core_utils/test_dd_tracing.py b/tests/unit/litellm_core_utils/test_dd_tracing.py index 30cae45e250..251f1842e6e 100644 --- a/tests/unit/litellm_core_utils/test_dd_tracing.py +++ b/tests/unit/litellm_core_utils/test_dd_tracing.py @@ -5,7 +5,7 @@ import pytest from litellm.litellm_core_utils.dd_tracing import ( - _should_use_dd_profiler, + should_use_dd_profiler, _should_use_dd_tracer, ) from litellm.litellm_core_utils.dd_tracing import tracer as dd_tracer @@ -87,7 +87,7 @@ def test_should_use_dd_profiler(): # Test when USE_DDPROFILER is True mock_get_secret.return_value = True - assert _should_use_dd_profiler() is True + assert should_use_dd_profiler() is True mock_get_secret.assert_called_once_with("USE_DDPROFILER", False) # Reset the mock for the next test @@ -95,5 +95,5 @@ def test_should_use_dd_profiler(): # Test when USE_DDPROFILER is False mock_get_secret.return_value = False - assert _should_use_dd_profiler() is False + assert should_use_dd_profiler() is False mock_get_secret.assert_called_once_with("USE_DDPROFILER", False) diff --git a/tests/unit/litellm_core_utils/test_get_litellm_params.py b/tests/unit/litellm_core_utils/test_get_litellm_params.py index c7c37abd46e..f21e475fd16 100644 --- a/tests/unit/litellm_core_utils/test_get_litellm_params.py +++ b/tests/unit/litellm_core_utils/test_get_litellm_params.py @@ -13,7 +13,7 @@ from litellm.constants import CONTROL_OPTIONS_KEY from litellm.litellm_core_utils.get_litellm_params import ( _OPTIONAL_KWARGS_KEYS, InvalidControlOption, - _get_base_model_from_litellm_call_metadata, + get_base_model_from_litellm_call_metadata, get_litellm_params, parse_control_options, stored_control_options, @@ -55,26 +55,31 @@ NAMED_PRICE_PARAMS: Final = frozenset( class TestGetBaseModelFromLitellmCallMetadata: def test_none_metadata_returns_none(self): - assert _get_base_model_from_litellm_call_metadata(None) is None + assert get_base_model_from_litellm_call_metadata(None) is None def test_empty_metadata_returns_none(self): - assert _get_base_model_from_litellm_call_metadata({}) is None + assert get_base_model_from_litellm_call_metadata({}) is None def test_missing_model_info_returns_none(self): - assert _get_base_model_from_litellm_call_metadata({"foo": "bar"}) is None + assert get_base_model_from_litellm_call_metadata({"foo": "bar"}) is None def test_model_info_none_returns_none(self): - assert _get_base_model_from_litellm_call_metadata({"model_info": None}) is None + assert get_base_model_from_litellm_call_metadata({"model_info": None}) is None def test_model_info_empty_dict_returns_none(self): - assert _get_base_model_from_litellm_call_metadata({"model_info": {}}) is None + assert get_base_model_from_litellm_call_metadata({"model_info": {}}) is None def test_returns_base_model(self): - result = _get_base_model_from_litellm_call_metadata( + result = get_base_model_from_litellm_call_metadata( {"model_info": {"base_model": "gpt-4"}} ) assert result == "gpt-4" + def test_returns_non_string_base_model_without_filtering(self): + result = get_base_model_from_litellm_call_metadata({"model_info": {"base_model": 123}}) + + assert result == 123 + class TestGetLitellmParamsKwargsExtraction: """Verify that optional kwargs are correctly extracted via sparse extraction.""" diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index e8a81d8ac87..f677589e1d6 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -147,7 +147,7 @@ async def test_ahealth_check_supports_image_edit_mode(): def test_update_model_params_with_health_check_tracking_information(): - """Test _update_model_params_with_health_check_tracking_information adds required tracking info.""" + """Test update_model_params_with_health_check_tracking_information adds required tracking info.""" initial_model_params = {"model": "gpt-3.5-turbo", "api_key": "test_key"} with patch( @@ -167,7 +167,7 @@ def test_update_model_params_with_health_check_tracking_information(): }, } - result = HealthCheckHelpers._update_model_params_with_health_check_tracking_information( + result = HealthCheckHelpers.update_model_params_with_health_check_tracking_information( initial_model_params ) @@ -445,7 +445,7 @@ async def test_realtime_health_check_uses_model_level_vertex_params(): ), patch.object( HealthCheckHelpers, - "_update_model_params_with_health_check_tracking_information", + "update_model_params_with_health_check_tracking_information", staticmethod(lambda model_params: model_params), ), ): diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index 5e85a5903c6..1c17f74ae7b 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -87,6 +87,35 @@ async def test_async_post_mcp_tool_call_hook_preserves_and_returns_content(loggi assert hooked_content.content == [TextContent(type="text", text="[REDACTED]")] +def test_response_cost_calculator_preserves_dynamic_call_details(logging_obj, monkeypatch): + cache_hit: Final = 1 + + class Provider: + value: Final = "custom" + + def startswith(self, prefix: str) -> bool: + return self.value.startswith(prefix) + + custom_llm_provider: Final = Provider() + + def response_cost_calculator(**kwargs: object) -> float: + assert kwargs["cache_hit"] is cache_hit + assert kwargs["custom_llm_provider"] is custom_llm_provider + return 0.0 + + logging_obj.model_call_details["cache_hit"] = cache_hit + logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider + logging_obj.optional_params = {} + monkeypatch.setattr(litellm, "response_cost_calculator", response_cost_calculator) + + result = ModelResponse( + model="gpt-4o-mini", + usage={"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + ) + + assert logging_obj.response_cost_calculator(result=result) == 0.0 + + @pytest.mark.asyncio async def test_async_post_mcp_tool_call_hook_chains_every_callback(logging_obj): from litellm.types.mcp import MCPPostCallResponseObject @@ -432,7 +461,7 @@ def test_use_custom_pricing_not_detected_litellm_metadata_no_pricing(): def test_response_cost_calculator_uses_router_model_id_from_litellm_metadata(): - """_response_cost_calculator should extract router_model_id from + """response_cost_calculator should extract router_model_id from litellm_params.litellm_metadata.model_info.id when the result object does not carry _hidden_params (e.g. ResponsesAPIResponse from /v1/responses streaming). Regression test for custom pricing on streaming responses.""" @@ -496,7 +525,7 @@ def test_response_cost_calculator_uses_router_model_id_from_litellm_metadata(): }, ) - cost = logging_obj._response_cost_calculator(result=response_obj) + cost = logging_obj.response_cost_calculator(result=response_obj) assert cost is not None, "Cost should not be None" expected_cost = (10 * custom_input_cost) + (5 * custom_output_cost) @@ -599,8 +628,8 @@ class TestZeroCostDiagnostic: logging_obj: Final = self._logging_obj(deployment_pricing) with caplog.at_level(logging.WARNING, logger="LiteLLM"): - first_cost: Final = logging_obj._response_cost_calculator(result=self._response(usage)) - second_cost: Final = logging_obj._response_cost_calculator(result=self._response(usage)) + first_cost: Final = logging_obj.response_cost_calculator(result=self._response(usage)) + second_cost: Final = logging_obj.response_cost_calculator(result=self._response(usage)) assert first_cost == 0.0 assert second_cost == 0.0 @@ -617,8 +646,8 @@ class TestZeroCostDiagnostic: logging_obj: Final = self._logging_obj(deployment_pricing, stream=True) with caplog.at_level(logging.WARNING, logger="LiteLLM"): - logging_obj._response_cost_calculator(result=self._response(usage=None)) - logging_obj._response_cost_calculator(result=self._response(usage)) + logging_obj.response_cost_calculator(result=self._response(usage=None)) + logging_obj.response_cost_calculator(result=self._response(usage)) if deployment_pricing is self.FREE_PRICING: assert logging_obj.model_call_details["zero_cost_diagnostic"] is None @@ -644,7 +673,7 @@ class TestZeroCostDiagnostic: ) with caplog.at_level(logging.WARNING, logger="LiteLLM"): - cost: Final = logging_obj._response_cost_calculator(result=event) + cost: Final = logging_obj.response_cost_calculator(result=event) assert cost == 0.0 if deployment_pricing is self.FREE_PRICING: @@ -711,7 +740,7 @@ class TestZeroCostDiagnostic: ) with caplog.at_level(logging.WARNING, logger="LiteLLM"): - cost: Final = logging_obj._response_cost_calculator( + cost: Final = logging_obj.response_cost_calculator( result=self._response(usage, model="lit7898-unmapped-model") ) @@ -726,7 +755,7 @@ class TestZeroCostDiagnostic: logging_obj: Final = self._logging_obj(deployment_pricing) with caplog.at_level(logging.WARNING, logger="LiteLLM"): - cost: Final = logging_obj._response_cost_calculator( + cost: Final = logging_obj.response_cost_calculator( result={"model": "gpt-5.4-nano", "usage": {"prompt_tokens": "n/a", "completion_tokens": 3}} ) @@ -741,10 +770,10 @@ class TestZeroCostDiagnostic: logging_obj: Final = self._logging_obj(deployment_pricing, stream=True, call_type="anthropic_messages") with caplog.at_level(logging.WARNING, logger="LiteLLM"): - logging_obj._response_cost_calculator(result=self._response(usage=None)) - logging_obj._response_cost_calculator(result=self._response(usage)) - logging_obj._response_cost_calculator(result=self._response(usage=None)) - logging_obj._response_cost_calculator(result=self._response(usage)) + logging_obj.response_cost_calculator(result=self._response(usage=None)) + logging_obj.response_cost_calculator(result=self._response(usage)) + logging_obj.response_cost_calculator(result=self._response(usage=None)) + logging_obj.response_cost_calculator(result=self._response(usage)) if deployment_pricing is self.FREE_PRICING: assert logging_obj.model_call_details["zero_cost_diagnostic"] is None @@ -765,15 +794,15 @@ class TestZeroCostDiagnostic: try: logging_obj: Final = self._logging_obj(self.QUERY_ONLY_PRICING) with caplog.at_level(logging.WARNING, logger="LiteLLM"): - assert logging_obj._response_cost_calculator(result=self._response(usage)) == 0.0 + assert logging_obj.response_cost_calculator(result=self._response(usage)) == 0.0 self._assert_flagged(logging_obj, caplog) self._route_to_deployment(logging_obj, priced_pricing, deployment_id=priced_id) - assert logging_obj._response_cost_calculator(result=self._response(usage)) == pytest.approx(5e-05) + assert logging_obj.response_cost_calculator(result=self._response(usage)) == pytest.approx(5e-05) assert logging_obj.model_call_details["zero_cost_diagnostic"] is None self._route_to_deployment(logging_obj, self.QUERY_ONLY_PRICING) - assert logging_obj._response_cost_calculator(result=self._response(usage)) == 0.0 + assert logging_obj.response_cost_calculator(result=self._response(usage)) == 0.0 assert logging_obj.model_call_details["zero_cost_diagnostic"]["reason"] == "missing_pricing_key" assert len(self._zero_cost_warnings(caplog)) == 1 @@ -796,8 +825,8 @@ class TestZeroCostDiagnostic: {}, model=f"openai/{requested_model}", deployment_id="lit7898-cost-map-deployment" ) with caplog.at_level(logging.WARNING, logger="LiteLLM"): - assert logging_obj._response_cost_calculator(result=self._response(usage, model=dated_model)) == 0.0 - assert logging_obj._response_cost_calculator(result=self._response(usage, model=requested_model)) == 0.0 + assert logging_obj.response_cost_calculator(result=self._response(usage, model=dated_model)) == 0.0 + assert logging_obj.response_cost_calculator(result=self._response(usage, model=requested_model)) == 0.0 assert logging_obj.model_call_details["zero_cost_diagnostic"]["pricing_model"] == requested_model warnings: Final = self._zero_cost_warnings(caplog) @@ -843,7 +872,7 @@ class TestZeroCostDiagnostic: logging_obj.model_call_details["cache_hit"] = True with caplog.at_level(logging.WARNING, logger="LiteLLM"): - assert logging_obj._response_cost_calculator(result=self._response(usage), cache_hit=False) == 0.0 + assert logging_obj.response_cost_calculator(result=self._response(usage), cache_hit=False) == 0.0 assert logging_obj.model_call_details["zero_cost_diagnostic"] is None assert self._zero_cost_warnings(caplog) == [] @@ -859,7 +888,7 @@ class TestZeroCostDiagnostic: response: Final = self._response(usage) response._response_ms = 1000.0 with caplog.at_level(logging.WARNING, logger="LiteLLM"): - assert logging_obj._response_cost_calculator(result=response) == pytest.approx(0.00042) + assert logging_obj.response_cost_calculator(result=response) == pytest.approx(0.00042) assert logging_obj.model_call_details["zero_cost_diagnostic"] is None assert self._zero_cost_warnings(caplog) == [] @@ -897,7 +926,7 @@ class TestZeroCostDiagnostic: usage, model=served_model, custom_llm_provider="azure", additional_headers=spillover_headers ) with caplog.at_level(logging.WARNING, logger="LiteLLM"): - assert logging_obj._response_cost_calculator(result=response) == 0.0 + assert logging_obj.response_cost_calculator(result=response) == 0.0 warnings: Final = self._zero_cost_warnings(caplog) if not spilled_over: @@ -1211,7 +1240,7 @@ class TestRetrieveBatchCostPassesModelIdentity: failed_requests=0, ) - monkeypatch.setattr(logging_module, "_handle_completed_batch", fake_handle_completed_batch) + monkeypatch.setattr(logging_module, "handle_completed_batch", fake_handle_completed_batch) obj = LitellmLogging( model="bedrock/global.anthropic.claude-sonnet-4-6", @@ -1242,7 +1271,7 @@ class TestRetrieveBatchCostPassesModelIdentity: finally: litellm.model_cost.pop(deployment_id, None) - assert captured, "_handle_completed_batch was never called" + assert captured, "handle_completed_batch was never called" assert captured["model_name"] == "bedrock/global.anthropic.claude-sonnet-4-6" assert captured["model_info"] is not None assert captured["model_info"]["input_cost_per_token"] == 0.0 @@ -1294,7 +1323,7 @@ class TestRetrieveBatchPricesOnlyFinalBatches: from litellm.litellm_core_utils import litellm_logging as logging_module handle_completed_batch = AsyncMock() - monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) + monkeypatch.setattr(logging_module, "handle_completed_batch", handle_completed_batch) batch = self._batch(status, output_file_id) await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) @@ -1317,7 +1346,7 @@ class TestRetrieveBatchPricesOnlyFinalBatches: failed_requests=0, ) ) - monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) + monkeypatch.setattr(logging_module, "handle_completed_batch", handle_completed_batch) batch = self._batch("completed", "file-out") await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) @@ -1880,7 +1909,7 @@ async def test_anthropic_messages_marks_litellm_params_async(): await asyncio.wait_for(logged.wait(), timeout=10) assert captured["litellm_params"].get("aanthropic_messages") is True - assert LitellmLogging._is_sync_litellm_request(captured["litellm_params"]) is False + assert LitellmLogging.is_sync_litellm_request(captured["litellm_params"]) is False logger.log_success_event.assert_not_called() finally: litellm.callbacks = original_callbacks @@ -1912,7 +1941,7 @@ async def test_arealtime_marks_litellm_params_async(monkeypatch): await asyncio.wait_for(async_logged.wait(), timeout=10) logger.log_failure_event.assert_not_called() assert captured["litellm_params"].get("_arealtime") is True - assert LitellmLogging._is_sync_litellm_request(captured["litellm_params"]) is False + assert LitellmLogging.is_sync_litellm_request(captured["litellm_params"]) is False @pytest.mark.asyncio @@ -1977,7 +2006,7 @@ async def test_agenerate_content_marks_litellm_params_async(): litellm_params = logging_obj.model_call_details.get("litellm_params", {}) assert litellm_params.get("agenerate_content") is True - assert LitellmLogging._is_sync_litellm_request(litellm_params) is False + assert LitellmLogging.is_sync_litellm_request(litellm_params) is False @pytest.mark.asyncio @@ -2234,14 +2263,14 @@ async def test_async_success_handler_runs_async_callbacks_in_the_post_response_p def test_is_sync_litellm_request(): - assert LitellmLogging._is_sync_litellm_request({}) is True - assert LitellmLogging._is_sync_litellm_request({"acompletion": True}) is False - assert LitellmLogging._is_sync_litellm_request({"allm_passthrough_route": True}) is False - assert LitellmLogging._is_sync_litellm_request({"_arealtime": True}) is False - assert LitellmLogging._is_sync_litellm_request({"aanthropic_messages": True}) is False - assert LitellmLogging._is_sync_litellm_request({"agenerate_content": True}) is False - assert LitellmLogging._is_sync_litellm_request({"agenerate_content_stream": True}) is False - assert LitellmLogging._is_sync_litellm_request({"aanthropic_messages": False}) is True + assert LitellmLogging.is_sync_litellm_request({}) is True + assert LitellmLogging.is_sync_litellm_request({"acompletion": True}) is False + assert LitellmLogging.is_sync_litellm_request({"allm_passthrough_route": True}) is False + assert LitellmLogging.is_sync_litellm_request({"_arealtime": True}) is False + assert LitellmLogging.is_sync_litellm_request({"aanthropic_messages": True}) is False + assert LitellmLogging.is_sync_litellm_request({"agenerate_content": True}) is False + assert LitellmLogging.is_sync_litellm_request({"agenerate_content_stream": True}) is False + assert LitellmLogging.is_sync_litellm_request({"aanthropic_messages": False}) is True def test_get_litellm_params_propagates_allm_passthrough_route(): @@ -2252,7 +2281,7 @@ def test_get_litellm_params_propagates_allm_passthrough_route(): params = get_litellm_params(allm_passthrough_route=True) assert params.get("allm_passthrough_route") is True - assert LitellmLogging._is_sync_litellm_request(params) is False + assert LitellmLogging.is_sync_litellm_request(params) is False @pytest.mark.asyncio @@ -2877,7 +2906,7 @@ def test_get_user_agent_tags(): def test_get_request_tags(): from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup - tags = StandardLoggingPayloadSetup._get_request_tags( + tags = StandardLoggingPayloadSetup.get_request_tags( litellm_params={"metadata": {"tags": ["test-tag"]}}, proxy_server_request={ "headers": { @@ -2891,6 +2920,17 @@ def test_get_request_tags(): assert "User-Agent: litellm/0.1.0" in tags +def test_get_request_tags_preserves_non_string_metadata_tags(): + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + tags = StandardLoggingPayloadSetup.get_request_tags( + litellm_params={"metadata": {"tags": [7]}}, + proxy_server_request={}, + ) + + assert 7 in tags + + def test_get_request_tags_from_metadata_and_litellm_metadata(): """ Test that _get_request_tags correctly picks tags from both 'metadata' and 'litellm_metadata'. @@ -2905,7 +2945,7 @@ def test_get_request_tags_from_metadata_and_litellm_metadata(): from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup # Test case 1: Tags in metadata only - tags = StandardLoggingPayloadSetup._get_request_tags( + tags = StandardLoggingPayloadSetup.get_request_tags( litellm_params={"metadata": {"tags": ["metadata-tag-1", "metadata-tag-2"]}}, proxy_server_request={}, ) @@ -2914,7 +2954,7 @@ def test_get_request_tags_from_metadata_and_litellm_metadata(): assert len([t for t in tags if not t.startswith("User-Agent:")]) == 2 # Test case 2: Tags in litellm_metadata only - tags = StandardLoggingPayloadSetup._get_request_tags( + tags = StandardLoggingPayloadSetup.get_request_tags( litellm_params={"litellm_metadata": {"tags": ["litellm-metadata-tag-1", "litellm-metadata-tag-2"]}}, proxy_server_request={}, ) @@ -2923,7 +2963,7 @@ def test_get_request_tags_from_metadata_and_litellm_metadata(): assert len([t for t in tags if not t.startswith("User-Agent:")]) == 2 # Test case 3: Tags in both - metadata should take priority - tags = StandardLoggingPayloadSetup._get_request_tags( + tags = StandardLoggingPayloadSetup.get_request_tags( litellm_params={ "metadata": {"tags": ["metadata-tag"]}, "litellm_metadata": {"tags": ["litellm-metadata-tag"]}, @@ -2935,14 +2975,14 @@ def test_get_request_tags_from_metadata_and_litellm_metadata(): assert len([t for t in tags if not t.startswith("User-Agent:")]) == 1 # Test case 4: No tags in either - tags = StandardLoggingPayloadSetup._get_request_tags( + tags = StandardLoggingPayloadSetup.get_request_tags( litellm_params={"metadata": {}, "litellm_metadata": {}}, proxy_server_request={}, ) assert len([t for t in tags if not t.startswith("User-Agent:")]) == 0 # Test case 5: None values for metadata/litellm_metadata - tags = StandardLoggingPayloadSetup._get_request_tags( + tags = StandardLoggingPayloadSetup.get_request_tags( litellm_params={"metadata": None, "litellm_metadata": None}, proxy_server_request={}, ) @@ -2950,7 +2990,7 @@ def test_get_request_tags_from_metadata_and_litellm_metadata(): assert len([t for t in tags if not t.startswith("User-Agent:")]) == 0 # Test case 6: Empty litellm_params - tags = StandardLoggingPayloadSetup._get_request_tags( + tags = StandardLoggingPayloadSetup.get_request_tags( litellm_params={}, proxy_server_request={}, ) @@ -2958,7 +2998,7 @@ def test_get_request_tags_from_metadata_and_litellm_metadata(): assert len([t for t in tags if not t.startswith("User-Agent:")]) == 0 # Test case 7: Metadata tags combined with user-agent tags - tags = StandardLoggingPayloadSetup._get_request_tags( + tags = StandardLoggingPayloadSetup.get_request_tags( litellm_params={"metadata": {"tags": ["custom-tag"]}}, proxy_server_request={ "headers": { @@ -2992,15 +3032,15 @@ def test_get_request_tags_does_not_mutate_original_tags(): } # Call _get_request_tags multiple times (simulating multiple callbacks) - tags1 = StandardLoggingPayloadSetup._get_request_tags( + tags1 = StandardLoggingPayloadSetup.get_request_tags( litellm_params=litellm_params, proxy_server_request=proxy_server_request, ) - tags2 = StandardLoggingPayloadSetup._get_request_tags( + tags2 = StandardLoggingPayloadSetup.get_request_tags( litellm_params=litellm_params, proxy_server_request=proxy_server_request, ) - tags3 = StandardLoggingPayloadSetup._get_request_tags( + tags3 = StandardLoggingPayloadSetup.get_request_tags( litellm_params=litellm_params, proxy_server_request=proxy_server_request, ) @@ -3150,7 +3190,7 @@ def test_response_cost_calculator_with_response_cost_in_hidden_params(logging_ob mock_response="Hello, world!", ) - response_cost = logging_obj._response_cost_calculator( + response_cost = logging_obj.response_cost_calculator( result=mock_response, ) @@ -3199,7 +3239,7 @@ def test_response_cost_calculator_native_generate_content_body_uses_usage_metada ) assert expected_cost > 0 - cost = logging_obj._response_cost_calculator(result=native_body) + cost = logging_obj.response_cost_calculator(result=native_body) assert cost == pytest.approx(expected_cost) @@ -3217,7 +3257,7 @@ def test_response_cost_calculator_does_not_transform_non_generate_content_dict() ) logging_obj.optional_params = {} - cost = logging_obj._response_cost_calculator( + cost = logging_obj.response_cost_calculator( result={"usageMetadata": {"promptTokenCount": 1000, "candidatesTokenCount": 500}} ) assert not cost @@ -3248,7 +3288,7 @@ def test_file_content_call_is_not_billed(call_type): """ result = HttpxBinaryResponseContent(httpx.Response(status_code=200, content=b"file contents")) - cost = _file_content_logging_obj(call_type)._response_cost_calculator(result=result) + cost = _file_content_logging_obj(call_type).response_cost_calculator(result=result) assert cost == 0.0 @@ -3271,7 +3311,7 @@ def test_speech_call_is_still_priced_from_input_characters(call_type): result = HttpxBinaryResponseContent(httpx.Response(status_code=200, content=b"audio bytes")) - cost = logging_obj._response_cost_calculator(result=result) + cost = logging_obj.response_cost_calculator(result=result) assert cost is not None assert cost > 0 @@ -3309,7 +3349,7 @@ def test_sentry_send_default_pii_opt_in(monkeypatch): def test_get_masked_values(): - from litellm.litellm_core_utils.litellm_logging import _get_masked_values + from litellm.litellm_core_utils.litellm_logging import get_masked_values sensitive_object = { "mode": "pre_call", @@ -3353,11 +3393,20 @@ def test_get_masked_values(): "presidio_anonymizer_api_base": None, "vertex_credentials": "{sensitive_api_key}", } - masked_values = _get_masked_values(sensitive_object, unmasked_length=4, number_of_asterisks=4) + masked_values = get_masked_values(sensitive_object, unmasked_length=4, number_of_asterisks=4) assert masked_values["presidio_anonymizer_api_base"] is None assert masked_values["vertex_credentials"] == "{s****y}" +def test_get_masked_values_keeps_original_non_string_key_behavior(): + from typing import cast + + from litellm.litellm_core_utils.litellm_logging import get_masked_values + + with pytest.raises(AttributeError): + get_masked_values({"api_key": cast(object, {1: "value"})}) + + @pytest.mark.asyncio async def test_e2e_generate_cold_storage_object_key_successful(): """ @@ -5286,7 +5335,7 @@ def test_success_handler_computes_cost_for_dict_response(): with ( patch.object( logging_obj, - "_response_cost_calculator", + "response_cost_calculator", return_value=expected_cost, ) as mock_calc, patch.object( @@ -5323,7 +5372,7 @@ def test_success_handler_preserves_precomputed_cost_for_dict_response(): with ( patch.object( logging_obj, - "_response_cost_calculator", + "response_cost_calculator", return_value=9.99, ) as mock_calc, patch.object( @@ -5362,7 +5411,7 @@ def test_success_handler_unified_helper_runs_for_typed_results(): with ( patch.object( logging_obj, - "_response_cost_calculator", + "response_cost_calculator", return_value=expected_cost, ) as mock_calc, patch.object( @@ -6675,7 +6724,7 @@ class TestNonInferenceCallTypesAreNotBilled: def test_creating_a_response_is_still_priced(self): """Guards the tests below: the same response object must cost money on the create path.""" - cost = self._logging_obj("aresponses")._response_cost_calculator(result=self._retrieved_response()) + cost = self._logging_obj("aresponses").response_cost_calculator(result=self._retrieved_response()) assert cost is not None and cost > 0 @pytest.mark.parametrize( @@ -6691,7 +6740,7 @@ class TestNonInferenceCallTypesAreNotBilled: ], ) def test_read_and_management_calls_cost_nothing(self, call_type): - cost = self._logging_obj(call_type)._response_cost_calculator(result=self._retrieved_response()) + cost = self._logging_obj(call_type).response_cost_calculator(result=self._retrieved_response()) assert cost == 0.0 def test_retrieved_usage_is_not_re_reported_in_standard_logging_payload(self): @@ -6728,7 +6777,7 @@ class TestNonInferenceCallTypesAreNotBilled: only billable usage. Zeroing it there means background jobs are never billed.""" cost = self._logging_obj( "aget_responses", litellm_metadata=self.BACKGROUND_POLL_METADATA - )._response_cost_calculator(result=self._retrieved_response()) + ).response_cost_calculator(result=self._retrieved_response()) assert cost is not None and cost > 0 def test_background_cost_poll_reports_usage_in_standard_logging_payload(self): @@ -6759,7 +6808,7 @@ class TestNonInferenceCallTypesAreNotBilled: def test_reading_a_background_response_is_still_priced(self): """A background create answers queued with no usage at all, so whoever reads the finished job is the first and only caller to see its tokens. Zeroing that read bills the job nothing.""" - cost = self._logging_obj("aget_responses")._response_cost_calculator( + cost = self._logging_obj("aget_responses").response_cost_calculator( result=self._retrieved_response(background=True) ) assert cost is not None and cost > 0 @@ -6792,7 +6841,7 @@ class TestNonInferenceCallTypesAreNotBilled: def test_reading_a_foreground_response_is_still_free(self): """Guards the test above against a blanket exemption: an explicit background=false read was already billed by its create and must stay at zero.""" - cost = self._logging_obj("aget_responses")._response_cost_calculator( + cost = self._logging_obj("aget_responses").response_cost_calculator( result=self._retrieved_response(background=False) ) assert cost == 0.0 @@ -7054,7 +7103,7 @@ async def test_streaming_success_callbacks_survive_cost_calculation_failure(): releasing.async_log_success_event = AsyncMock() patcher, logging_obj = _streaming_logging_obj_with_callbacks([releasing]) - with patcher, patch.object(logging_obj, "_response_cost_calculator", side_effect=ValueError("bad usage block")): + with patcher, patch.object(logging_obj, "response_cost_calculator", side_effect=ValueError("bad usage block")): await logging_obj.async_success_handler(result=_assembled_stream_result()) assert logging_obj.model_call_details["response_cost"] is None @@ -7218,7 +7267,7 @@ def test_response_cost_calculator_prices_mantle_calls_on_the_served_region(monke choices=[{"message": {"role": "assistant", "content": "hello"}, "index": 0, "finish_reason": "stop"}], usage={"prompt_tokens": 38, "completion_tokens": 20, "total_tokens": 58}, ) - return logging_obj._response_cost_calculator(result=response) + return logging_obj.response_cost_calculator(result=response) commercial = litellm.model_cost["bedrock_mantle/xai.grok-4.3"] gov = litellm.model_cost["bedrock_mantle/us-gov-west-1/xai.grok-4.3"] @@ -7272,7 +7321,7 @@ def test_response_cost_calculator_prices_proxy_vertex_calls_on_the_configured_lo choices=[{"message": {"role": "assistant", "content": "hello"}, "index": 0, "finish_reason": "stop"}], usage={"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, ) - return logging_obj._response_cost_calculator(result=response) + return logging_obj.response_cost_calculator(result=response) info = litellm.model_cost["vertex_ai/gemini-3.5-flash"] expected_global = 10 * info["input_cost_per_token"] + 5 * info["output_cost_per_token"] @@ -7331,7 +7380,7 @@ def test_response_cost_calculator_prices_proxy_vertex_image_calls_on_the_configu total_tokens=1220, ), ) - return logging_obj._response_cost_calculator(result=response) + return logging_obj.response_cost_calculator(result=response) expected_global = 100 * 5e-07 + 1120 * 6e-05 @@ -8794,7 +8843,7 @@ class TestAzurePTUSpilloverCost: response = self._response() response._hidden_params["additional_headers"] = {"llm_provider-x-ms-is-spilled-over": "true"} - assert obj._response_cost_calculator(result=response) == pytest.approx(self.EXPECTED_SPILL_COST) + assert obj.response_cost_calculator(result=response) == pytest.approx(self.EXPECTED_SPILL_COST) finally: self._unregister_models() @@ -8807,7 +8856,7 @@ class TestAzurePTUSpilloverCost: "x-ms-spillover-from-deployment": "ptu-dep", } - assert obj._response_cost_calculator(result=self._response()) == pytest.approx(self.EXPECTED_SPILL_COST) + assert obj.response_cost_calculator(result=self._response()) == pytest.approx(self.EXPECTED_SPILL_COST) finally: self._unregister_models() @@ -8816,7 +8865,7 @@ class TestAzurePTUSpilloverCost: try: obj = self._logging_obj(dict(self.PTU_MODEL_INFO), flag="True", litellm_rate=0.0, monkeypatch=monkeypatch) - assert obj._response_cost_calculator(result=self._response()) == 0.0 + assert obj.response_cost_calculator(result=self._response()) == 0.0 finally: self._unregister_models() @@ -8827,7 +8876,7 @@ class TestAzurePTUSpilloverCost: response = self._response() response._hidden_params["additional_headers"] = {"llm_provider-x-ms-is-spilled-over": "true"} - assert obj._response_cost_calculator(result=response) == 0.0 + assert obj.response_cost_calculator(result=response) == 0.0 finally: self._unregister_models() @@ -8846,7 +8895,7 @@ class TestAzurePTUSpilloverCost: response = self._response() response._hidden_params["additional_headers"] = {"llm_provider-x-ms-is-spilled-over": "true"} - assert obj._response_cost_calculator(result=response) == pytest.approx(150 * 1e-6) + assert obj.response_cost_calculator(result=response) == pytest.approx(150 * 1e-6) finally: litellm.model_cost.pop(custom_model_id, None) self._unregister_models() @@ -8892,7 +8941,7 @@ def test_get_assembled_streaming_response_bills_a_provider_reported_usage_cost() ) assert assembled._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] == 0.0042 - assert logging_obj._response_cost_calculator(result=assembled) == 0.0042 + assert logging_obj.response_cost_calculator(result=assembled) == 0.0042 @@ -8910,8 +8959,8 @@ def test_response_cost_calculator_prices_terminal_responses_event_from_its_respo ) event: Final = ResponseCompletedEvent(type="response.completed", response=inner_response) - event_cost: Final = logging_obj._response_cost_calculator(result=event) - inner_cost: Final = logging_obj._response_cost_calculator(result=inner_response) + event_cost: Final = logging_obj.response_cost_calculator(result=event) + inner_cost: Final = logging_obj.response_cost_calculator(result=inner_response) assert event_cost is not None and event_cost > 0 assert event_cost == inner_cost @@ -9137,7 +9186,7 @@ def test_responses_completed_event_bills_the_served_service_tier(): ) event: Final = ResponseCompletedEvent(type="response.completed", response=inner) - cost: Final = logging_obj._response_cost_calculator(result=event) # pyright: ignore[reportPrivateUsage] # parity with the suite's own direct calls + cost: Final = logging_obj.response_cost_calculator(result=event) billed_response: Final = ModelResponse( model="gpt-5.1", diff --git a/tests/unit/litellm_core_utils/test_llm_request_utils.py b/tests/unit/litellm_core_utils/test_llm_request_utils.py index 765c47547ce..350419b49a5 100644 --- a/tests/unit/litellm_core_utils/test_llm_request_utils.py +++ b/tests/unit/litellm_core_utils/test_llm_request_utils.py @@ -2,6 +2,7 @@ import httpx import pytest from litellm.litellm_core_utils.llm_request_utils import ( + ensure_extra_body_is_safe, flatten_form_field_values, serialize_multipart_form_fields, ) @@ -105,3 +106,30 @@ def test_flatten_form_field_values_rejects_over_deep_nesting(): assert isinstance(nested, dict) with pytest.raises(ValueError, match="max depth"): flatten_form_field_values(nested) + + +def test_ensure_extra_body_is_safe_converts_prompt_without_filtering_metadata_keys(): + from typing import cast + + class Prompt: + pass + + prompt = Prompt() + metadata: dict[object, object] = {"prompt": prompt, 1: "retained"} + extra_body = cast(dict[str, object], {"metadata": metadata}) + + result = ensure_extra_body_is_safe(extra_body) + + assert result is extra_body + assert metadata == {"prompt": prompt.__dict__, 1: "retained"} + + +def test_ensure_extra_body_is_safe_returns_non_dict_unchanged(): + from collections import UserDict + from typing import cast + + extra_body = UserDict({"metadata": {"prompt": object()}}) + + result = ensure_extra_body_is_safe(cast(dict[str, object], extra_body)) + + assert result is extra_body diff --git a/tests/unit/litellm_core_utils/test_model_param_helper.py b/tests/unit/litellm_core_utils/test_model_param_helper.py index 2c45b333817..2da324158cf 100644 --- a/tests/unit/litellm_core_utils/test_model_param_helper.py +++ b/tests/unit/litellm_core_utils/test_model_param_helper.py @@ -5,8 +5,8 @@ from litellm.litellm_core_utils.model_param_helper import ModelParamHelper def test_get_all_llm_api_params_is_correct(): """The cached result must equal a fresh, uncached computation.""" - cached = ModelParamHelper._get_all_llm_api_params() - uncached = ModelParamHelper._get_all_llm_api_params.__wrapped__() + cached = ModelParamHelper.get_all_llm_api_params() + uncached = ModelParamHelper.get_all_llm_api_params.__wrapped__() assert cached == uncached assert {"model", "temperature", "stream"} <= cached assert "metadata" not in cached # excluded via _get_exclude_kwargs @@ -16,6 +16,6 @@ def test_get_all_llm_api_params_is_memoized(): """Regression: the param set is static and is rebuilt on every request via the cache-key and spend-logging paths, so it must be memoized. Without the cache each call returns a freshly built set (a different object).""" - first = ModelParamHelper._get_all_llm_api_params() - second = ModelParamHelper._get_all_llm_api_params() + first = ModelParamHelper.get_all_llm_api_params() + second = ModelParamHelper.get_all_llm_api_params() assert first is second diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index 7e6d4d24905..a7a7ee77b9e 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -2022,7 +2022,7 @@ async def test_provider_path_suppresses_duplicate_session_created_after_syntheti model="gemini-2.5-flash", ) # Simulate synthetic session.created already sent by llm_http_handler. - streaming._session_created_sent_to_client = True + streaming.session_created_sent_to_client = True await streaming.backend_to_client_send_messages() @@ -2074,7 +2074,7 @@ async def test_duplicate_session_created_still_triggers_guardrail_turn_detection model="gemini-2.5-flash", ) # Synthetic session.created already sent by llm_http_handler. - streaming._session_created_sent_to_client = True + streaming.session_created_sent_to_client = True streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] streaming._send_to_backend = AsyncMock() # type: ignore[method-assign] @@ -2897,7 +2897,7 @@ def test_setup_folds_in_auto_response_disable_when_transcription_guardrail_activ } } ) - out = json.loads(streaming._maybe_inject_guardrail_auto_response_disable(setup)) + out = json.loads(streaming.maybe_inject_guardrail_auto_response_disable(setup)) aad = out["setup"]["realtimeInputConfig"]["automaticActivityDetection"] assert aad["disabled"] is True @@ -2908,7 +2908,7 @@ def test_setup_unchanged_without_transcription_guardrail(monkeypatch: pytest.Mon monkeypatch.setattr(litellm, "callbacks", []) streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) setup = json.dumps({"setup": {"model": "x", "generationConfig": {"responseModalities": ["AUDIO"]}}}) - out = streaming._maybe_inject_guardrail_auto_response_disable(setup) + out = streaming.maybe_inject_guardrail_auto_response_disable(setup) assert json.loads(out) == json.loads(setup) @@ -2920,7 +2920,7 @@ def test_non_bidi_setup_left_untouched_for_followup_capable_providers(monkeypatc monkeypatch.setattr(litellm, "callbacks", [_transcription_guardrail()]) streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) msg = json.dumps({"type": "session.update", "session": {"instructions": "hi"}}) - assert streaming._maybe_inject_guardrail_auto_response_disable(msg) == msg + assert streaming.maybe_inject_guardrail_auto_response_disable(msg) == msg @pytest.mark.asyncio diff --git a/tests/unit/litellm_core_utils/test_sensitive_data_masker.py b/tests/unit/litellm_core_utils/test_sensitive_data_masker.py index 5f73f3863e6..2182620164c 100644 --- a/tests/unit/litellm_core_utils/test_sensitive_data_masker.py +++ b/tests/unit/litellm_core_utils/test_sensitive_data_masker.py @@ -134,11 +134,11 @@ def test_short_secrets_are_fully_masked(): masker = SensitiveDataMasker() # Boundary: exactly 8 chars previously returned verbatim. - assert masker._mask_value("abcd1234") == "********" + assert masker.mask_value("abcd1234") == "********" # Below threshold previously hit the early return and leaked verbatim. - assert masker._mask_value("sk-12") == "*****" + assert masker.mask_value("sk-12") == "*****" # Values above the threshold must still partially reveal, not over-mask. - assert masker._mask_value("abcd12345") == "abcd*2345" + assert masker.mask_value("abcd12345") == "abcd*2345" masked = masker.mask_dict({"redis_password": "pass1234", "api_key": "sk-7a"}) assert masked["redis_password"] == "********" @@ -155,10 +155,10 @@ def test_mask_short_values_false_keeps_short_values_readable(): masker = SensitiveDataMasker(visible_prefix=50, visible_suffix=0, mask_short_values=False) short = "Test exception for structure validation" - assert masker._mask_value(short) == short + assert masker.mask_value(short) == short long_value = "x" * 60 - masked = masker._mask_value(long_value) + masked = masker.mask_value(long_value) assert masked.startswith("x" * 50) assert masked.endswith("*" * 10) assert len(masked) == 60 diff --git a/tests/unit/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py index f88e082d577..d4c6a781498 100644 --- a/tests/unit/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -405,22 +405,22 @@ def test_multi_chunk_reasoning_and_content( def test_strip_sse_data_from_chunk(): """Test the static method that strips 'data: ' prefix from SSE chunks""" # Test with string inputs - assert CustomStreamWrapper._strip_sse_data_from_chunk("data: content") == "content" + assert CustomStreamWrapper.strip_sse_data_from_chunk("data: content") == "content" assert ( - CustomStreamWrapper._strip_sse_data_from_chunk("data: spaced content") + CustomStreamWrapper.strip_sse_data_from_chunk("data: spaced content") == " spaced content" ) assert ( - CustomStreamWrapper._strip_sse_data_from_chunk("regular content") + CustomStreamWrapper.strip_sse_data_from_chunk("regular content") == "regular content" ) assert ( - CustomStreamWrapper._strip_sse_data_from_chunk("regular content with data:") + CustomStreamWrapper.strip_sse_data_from_chunk("regular content with data:") == "regular content with data:" ) # Test with None input - assert CustomStreamWrapper._strip_sse_data_from_chunk(None) is None + assert CustomStreamWrapper.strip_sse_data_from_chunk(None) is None @pytest.mark.parametrize("sync_mode", [True, False]) diff --git a/tests/unit/litellm_core_utils/test_streaming_overhead.py b/tests/unit/litellm_core_utils/test_streaming_overhead.py index 8fb0659ab5a..7eb1c467fb6 100644 --- a/tests/unit/litellm_core_utils/test_streaming_overhead.py +++ b/tests/unit/litellm_core_utils/test_streaming_overhead.py @@ -39,7 +39,7 @@ def _make_logging_obj(provider: str = "anthropic") -> MagicMock: logging_obj.stream_options = None logging_obj.messages = [{"role": "user", "content": "hi"}] logging_obj.completion_start_time = None - logging_obj._llm_caching_handler = None + logging_obj.llm_caching_handler = None return logging_obj diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index e0c5c22d420..163897c2609 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -1635,6 +1635,14 @@ def test_token_counter_uses_the_tokenizer_of_each_model_family_and_of_a_custom_t } +def test_empty_custom_tokenizer_uses_model_tokenizer() -> None: + text_value: Final = "A tokenizer fallback should preserve the model encoding." + custom_count: Final = _get_exact_count_function("gpt-3.5-turbo", {})(text_value) + model_count: Final = _get_exact_count_function("gpt-3.5-turbo", None)(text_value) + + assert custom_count == model_count + + def _threshold_test_messages(turns: int) -> list[dict]: messages: list[dict] = [{"role": "system", "content": "You are a terse assistant. " * 20}] for index in range(turns): diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py b/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py index fc9d77a10a2..faa58fa7f07 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_mcp_handler.py @@ -190,12 +190,12 @@ async def test_anthropic_messages_with_mcp_forwards_the_callers_mcp_credentials( else __import__( "litellm.responses.mcp.litellm_proxy_mcp_handler", fromlist=["LiteLLM_Proxy_MCP_Handler"] ).LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", new=process, ), patch.object( import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", new=execute, ), patch("litellm.anthropic_messages", new=AsyncMock(side_effect=responses)), @@ -272,12 +272,12 @@ async def test_anthropic_messages_with_mcp_hands_execution_the_requests_served_t patch.object(MCPRequestContext, "resolve", return_value=MCPRequestContext(user_api_key_auth="auth")), patch.object( import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", new=process, ), patch.object( import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", new=execute, ), patch("litellm.anthropic_messages", new=AsyncMock(side_effect=responses)), @@ -323,12 +323,12 @@ async def test_anthropic_messages_with_mcp_stops_when_every_tool_call_is_skipped patch.object(MCPRequestContext, "resolve", return_value=MCPRequestContext(user_api_key_auth="auth")), patch.object( import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", new=AsyncMock(return_value=([], {})), ), patch.object( import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", new=AsyncMock(return_value=[]), ), patch("litellm.anthropic_messages", new=anthropic_messages_mock), diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py b/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py index 54e25da6ab9..e6963c33a56 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py @@ -615,7 +615,7 @@ async def test_async_sse_wrapper_reraises_upstream_error_to_connected_client(): request_body={}, ) detached_hook = _DetachedFailureRecorder() - iterator.litellm_logging_obj._on_detached_stream_failure = detached_hook + iterator.litellm_logging_obj.on_detached_stream_failure = detached_hook received = [] @@ -660,7 +660,7 @@ async def test_async_sse_wrapper_logs_failure_on_upstream_error_after_disconnect request_body={}, ) detached_hook = _DetachedFailureRecorder() - iterator.litellm_logging_obj._on_detached_stream_failure = detached_hook + iterator.litellm_logging_obj.on_detached_stream_failure = detached_hook gen = iterator.async_sse_wrapper(_gated_failing_stream()) received = [await gen.__anext__(), await gen.__anext__()] @@ -701,7 +701,7 @@ async def test_async_sse_wrapper_logs_failure_when_queued_error_is_never_consume request_body={}, ) detached_hook = _DetachedFailureRecorder() - iterator.litellm_logging_obj._on_detached_stream_failure = detached_hook + iterator.litellm_logging_obj.on_detached_stream_failure = detached_hook gen = iterator.async_sse_wrapper(_failing_stream()) received = [await gen.__anext__(), await gen.__anext__()] diff --git a/tests/unit/llms/azure/passthrough/test_azure_passthrough_transformation.py b/tests/unit/llms/azure/passthrough/test_azure_passthrough_transformation.py index 0fcc9ef0034..bc8c1000fa4 100644 --- a/tests/unit/llms/azure/passthrough/test_azure_passthrough_transformation.py +++ b/tests/unit/llms/azure/passthrough/test_azure_passthrough_transformation.py @@ -158,7 +158,7 @@ def test_azure_passthrough_embeddings_relay_is_costed_per_input_token(): assert isinstance(result, EmbeddingResponse) assert logging_obj.call_type == "aembedding" assert per_token > 0 - assert logging_obj._response_cost_calculator(result=result) == pytest.approx(1000 * per_token) + assert logging_obj.response_cost_calculator(result=result) == pytest.approx(1000 * per_token) def test_azure_passthrough_responses_relay_is_costed_per_token(): @@ -167,7 +167,7 @@ def test_azure_passthrough_responses_relay_is_costed_per_token(): assert isinstance(result, ResponsesAPIResponse) assert logging_obj.call_type == "aresponses" - assert logging_obj._response_cost_calculator(result=result) == pytest.approx( + assert logging_obj.response_cost_calculator(result=result) == pytest.approx( 1000 * info["input_cost_per_token"] + 100 * info["output_cost_per_token"] ) @@ -321,7 +321,7 @@ def test_azure_passthrough_streaming_responses_chunks_are_costed_per_token(): assert isinstance(response, ResponseCompletedEvent) assert response.response.usage.input_tokens == 1000 assert logging_obj.call_type == "aresponses" - assert logging_obj._response_cost_calculator(result=response.response) == pytest.approx( + assert logging_obj.response_cost_calculator(result=response.response) == pytest.approx( 1000 * info["input_cost_per_token"] + 100 * info["output_cost_per_token"] ) diff --git a/tests/unit/llms/azure/test_azure_common_utils.py b/tests/unit/llms/azure/test_azure_common_utils.py index 5f254a3bc6c..302175c46f8 100644 --- a/tests/unit/llms/azure/test_azure_common_utils.py +++ b/tests/unit/llms/azure/test_azure_common_utils.py @@ -309,7 +309,7 @@ def test_initialize_with_ad_token_provider(setup_mocks, monkeypatch): def test_initialize_with_enable_token_refresh(setup_mocks, monkeypatch): - litellm._turn_on_debug() + litellm.turn_on_debug() # Enable token refresh monkeypatch.delenv("AZURE_CLIENT_ID", raising=False) monkeypatch.delenv("AZURE_CLIENT_SECRET", raising=False) diff --git a/tests/unit/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py b/tests/unit/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py index 7c48512adb8..24f21751065 100644 --- a/tests/unit/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py +++ b/tests/unit/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py @@ -419,7 +419,7 @@ def test_ocr_result_routes_the_relay_to_per_page_ocr_costing(): assert isinstance(result, OCRResponse) assert logging_obj.call_type == "aocr" assert per_page > 0 - assert logging_obj._response_cost_calculator(result=result) == pytest.approx(2 * per_page) + assert logging_obj.response_cost_calculator(result=result) == pytest.approx(2 * per_page) def test_ocr_binding_receives_the_endpoint_without_the_model_segment(): @@ -504,7 +504,7 @@ def test_foundry_embeddings_relay_is_costed_per_input_token(): assert isinstance(result, EmbeddingResponse) assert logging_obj.call_type == "aembedding" assert per_token > 0 - assert logging_obj._response_cost_calculator(result=result) == pytest.approx(1200 * per_token) + assert logging_obj.response_cost_calculator(result=result) == pytest.approx(1200 * per_token) def test_cohere_rerank_relay_is_costed_per_search_unit(): @@ -516,7 +516,7 @@ def test_cohere_rerank_relay_is_costed_per_search_unit(): assert isinstance(result, RerankResponse) assert logging_obj.call_type == "arerank" assert per_query > 0 - assert logging_obj._response_cost_calculator(result=result) == pytest.approx(2 * per_query) + assert logging_obj.response_cost_calculator(result=result) == pytest.approx(2 * per_query) def test_image_generation_relay_is_costed_per_image(): @@ -528,7 +528,7 @@ def test_image_generation_relay_is_costed_per_image(): assert isinstance(result, ImageResponse) assert logging_obj.call_type == "aimage_generation" assert per_image > 0 - assert logging_obj._response_cost_calculator(result=result) == pytest.approx(per_image) + assert logging_obj.response_cost_calculator(result=result) == pytest.approx(per_image) def test_flux_2_relay_through_the_provider_route_is_costed_per_image(): @@ -539,7 +539,7 @@ def test_flux_2_relay_through_the_provider_route_is_costed_per_image(): assert isinstance(result, ImageResponse) assert logging_obj.call_type == "aimage_generation" - assert logging_obj._response_cost_calculator(result=result) == pytest.approx(per_image) + assert logging_obj.response_cost_calculator(result=result) == pytest.approx(per_image) def test_rejected_rerank_relay_keeps_the_passthrough_object_and_call_type(): @@ -598,7 +598,7 @@ def test_streaming_responses_chunks_through_a_router_relay_are_costed_like_azure assert response is not None assert response.response.usage.output_tokens == 100 assert logging_obj.call_type == "aresponses" - assert logging_obj._response_cost_calculator(result=response.response) == pytest.approx( + assert logging_obj.response_cost_calculator(result=response.response) == pytest.approx( 1000 * info["input_cost_per_token"] + 100 * info["output_cost_per_token"] ) diff --git a/tests/unit/llms/bedrock/chat/test_streaming_choice_index.py b/tests/unit/llms/bedrock/chat/test_streaming_choice_index.py index 2364b7dd145..8f0743555f4 100644 --- a/tests/unit/llms/bedrock/chat/test_streaming_choice_index.py +++ b/tests/unit/llms/bedrock/chat/test_streaming_choice_index.py @@ -1,113 +1,113 @@ -""" -Test that Bedrock streaming responses always use choice index 0, -regardless of contentBlockIndex value. - -Bedrock's contentBlockIndex identifies content blocks within a message (e.g., -text=0, toolUse=1), NOT parallel completions. Since Bedrock doesn't support -n > 1, all chunks must use choice index 0. - -References: -- Bedrock InferenceConfiguration (no n parameter): - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_InferenceConfiguration.html -- OpenAI choice.index (for n > 1): - https://platform.openai.com/docs/api-reference/chat/object -""" - -from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder - - -class TestBedrockStreamingChoiceIndex: - """Test that all streaming chunks use choice index 0.""" - - def test_tool_call_chunk_uses_choice_index_zero(self): - """ - Core regression test: tool call chunks must use choice index 0, - not contentBlockIndex (which is 1 for tool calls). - - This was the bug - contentBlockIndex was incorrectly used as choice.index, - breaking OpenAI SDK's ChatCompletionAccumulator. - """ - handler = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") - - # First, simulate a tool use start event on contentBlockIndex 1 - start_chunk = { - "start": { - "toolUse": { - "toolUseId": "tooluse_abc123", - "name": "get_weather", - } - }, - "contentBlockIndex": 1, # Tool calls are on index 1 - } - - start_result = handler.converse_chunk_parser(start_chunk) - - # Choice index should be 0, NOT contentBlockIndex (1) - assert start_result.choices[0].index == 0 - assert start_result.choices[0].delta.tool_calls is not None - assert start_result.choices[0].delta.tool_calls[0]["id"] == "tooluse_abc123" - - # Now simulate tool use delta on contentBlockIndex 1 - delta_chunk = { - "delta": {"toolUse": {"input": '{"location": "San Francisco"}'}}, - "contentBlockIndex": 1, # Tool calls are on index 1 - } - - delta_result = handler.converse_chunk_parser(delta_chunk) - - # Choice index should still be 0, NOT contentBlockIndex (1) - assert delta_result.choices[0].index == 0 - assert delta_result.choices[0].delta.tool_calls is not None - assert ( - delta_result.choices[0].delta.tool_calls[0]["function"]["arguments"] - == '{"location": "San Francisco"}' - ) - - def test_mixed_content_blocks_all_use_choice_index_zero(self): - """ - Integration test simulating a realistic streaming session: - text (contentBlockIndex=0) → tool call (contentBlockIndex=1) → finish. - - All chunks must have choice.index=0 for OpenAI SDK compatibility. - """ - handler = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") - - # Chunk 1: Text on contentBlockIndex 0 - text_chunk = { - "delta": {"text": "Let me check the weather."}, - "contentBlockIndex": 0, - } - result1 = handler.converse_chunk_parser(text_chunk) - assert result1.choices[0].index == 0, "Text chunk should have index=0" - - # Chunk 2: Tool call start on contentBlockIndex 1 - tool_start_chunk = { - "start": { - "toolUse": { - "toolUseId": "tool_xyz", - "name": "get_weather", - } - }, - "contentBlockIndex": 1, - } - result2 = handler.converse_chunk_parser(tool_start_chunk) - assert ( - result2.choices[0].index == 0 - ), "Tool start should have index=0, not contentBlockIndex=1" - - # Chunk 3: Tool call delta on contentBlockIndex 1 - tool_delta_chunk = { - "delta": {"toolUse": {"input": '{"city": "NYC"}'}}, - "contentBlockIndex": 1, - } - result3 = handler.converse_chunk_parser(tool_delta_chunk) - assert ( - result3.choices[0].index == 0 - ), "Tool delta should have index=0, not contentBlockIndex=1" - - # Chunk 4: Finish reason - finish_chunk = { - "stopReason": "tool_use", - } - result4 = handler.converse_chunk_parser(finish_chunk) - assert result4.choices[0].index == 0, "Finish reason should have index=0" +""" +Test that Bedrock streaming responses always use choice index 0, +regardless of contentBlockIndex value. + +Bedrock's contentBlockIndex identifies content blocks within a message (e.g., +text=0, toolUse=1), NOT parallel completions. Since Bedrock doesn't support +n > 1, all chunks must use choice index 0. + +References: +- Bedrock InferenceConfiguration (no n parameter): + https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_InferenceConfiguration.html +- OpenAI choice.index (for n > 1): + https://platform.openai.com/docs/api-reference/chat/object +""" + +from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder + + +class TestBedrockStreamingChoiceIndex: + """Test that all streaming chunks use choice index 0.""" + + def test_tool_call_chunk_uses_choice_index_zero(self): + """ + Core regression test: tool call chunks must use choice index 0, + not contentBlockIndex (which is 1 for tool calls). + + This was the bug - contentBlockIndex was incorrectly used as choice.index, + breaking OpenAI SDK's ChatCompletionAccumulator. + """ + handler = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") + + # First, simulate a tool use start event on contentBlockIndex 1 + start_chunk = { + "start": { + "toolUse": { + "toolUseId": "tooluse_abc123", + "name": "get_weather", + } + }, + "contentBlockIndex": 1, # Tool calls are on index 1 + } + + start_result = handler.converse_chunk_parser(start_chunk) + + # Choice index should be 0, NOT contentBlockIndex (1) + assert start_result.choices[0].index == 0 + assert start_result.choices[0].delta.tool_calls is not None + assert start_result.choices[0].delta.tool_calls[0]["id"] == "tooluse_abc123" + + # Now simulate tool use delta on contentBlockIndex 1 + delta_chunk = { + "delta": {"toolUse": {"input": '{"location": "San Francisco"}'}}, + "contentBlockIndex": 1, # Tool calls are on index 1 + } + + delta_result = handler.converse_chunk_parser(delta_chunk) + + # Choice index should still be 0, NOT contentBlockIndex (1) + assert delta_result.choices[0].index == 0 + assert delta_result.choices[0].delta.tool_calls is not None + assert ( + delta_result.choices[0].delta.tool_calls[0]["function"]["arguments"] + == '{"location": "San Francisco"}' + ) + + def test_mixed_content_blocks_all_use_choice_index_zero(self): + """ + Integration test simulating a realistic streaming session: + text (contentBlockIndex=0) → tool call (contentBlockIndex=1) → finish. + + All chunks must have choice.index=0 for OpenAI SDK compatibility. + """ + handler = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") + + # Chunk 1: Text on contentBlockIndex 0 + text_chunk = { + "delta": {"text": "Let me check the weather."}, + "contentBlockIndex": 0, + } + result1 = handler.converse_chunk_parser(text_chunk) + assert result1.choices[0].index == 0, "Text chunk should have index=0" + + # Chunk 2: Tool call start on contentBlockIndex 1 + tool_start_chunk = { + "start": { + "toolUse": { + "toolUseId": "tool_xyz", + "name": "get_weather", + } + }, + "contentBlockIndex": 1, + } + result2 = handler.converse_chunk_parser(tool_start_chunk) + assert ( + result2.choices[0].index == 0 + ), "Tool start should have index=0, not contentBlockIndex=1" + + # Chunk 3: Tool call delta on contentBlockIndex 1 + tool_delta_chunk = { + "delta": {"toolUse": {"input": '{"city": "NYC"}'}}, + "contentBlockIndex": 1, + } + result3 = handler.converse_chunk_parser(tool_delta_chunk) + assert ( + result3.choices[0].index == 0 + ), "Tool delta should have index=0, not contentBlockIndex=1" + + # Chunk 4: Finish reason + finish_chunk = { + "stopReason": "tool_use", + } + result4 = handler.converse_chunk_parser(finish_chunk) + assert result4.choices[0].index == 0, "Finish reason should have index=0" diff --git a/tests/unit/llms/bedrock/realtime/test_bedrock_realtime_handler.py b/tests/unit/llms/bedrock/realtime/test_bedrock_realtime_handler.py index 3aa827beb80..b1836cb6c46 100644 --- a/tests/unit/llms/bedrock/realtime/test_bedrock_realtime_handler.py +++ b/tests/unit/llms/bedrock/realtime/test_bedrock_realtime_handler.py @@ -709,7 +709,7 @@ class TestBedrockRealtimeProviderFailurePropagation: ) assert replay.value.status_code == 400, "a committed session must not be silently restarted on a fallback" - assert not litellm._should_retry(replay.value.status_code), "the router must not retry the replay refusal" + assert not litellm.should_retry(replay.value.status_code), "the router must not retry the replay refusal" assert "Nova Sonic stream broke" in replay.value.message, "the router surfaces the last attempt's error" @pytest.mark.asyncio diff --git a/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py b/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py index b8380b7adb4..4d5990d8d18 100644 --- a/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py +++ b/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py @@ -363,7 +363,7 @@ class TestGithubCopilotResponsesAPIRouting: in the (already-merged) model info; otherwise returns None so the dispatcher routes through the chat-completions translation bridge.""" - @patch("litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper") + @patch("litellm.llms.github_copilot.responses.transformation.cached_get_model_info_helper") def test_returns_config_when_mode_is_responses(self, mock_get_info): """``mode=responses`` returns native config.""" mock_get_info.return_value = {"mode": "responses"} @@ -373,7 +373,7 @@ class TestGithubCopilotResponsesAPIRouting: ) assert isinstance(config, GithubCopilotResponsesAPIConfig) - @patch("litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper") + @patch("litellm.llms.github_copilot.responses.transformation.cached_get_model_info_helper") def test_returns_none_when_mode_is_chat(self, mock_get_info): """``mode=chat`` returns None so dispatcher uses bridge.""" mock_get_info.return_value = {"mode": "chat"} @@ -383,7 +383,7 @@ class TestGithubCopilotResponsesAPIRouting: ) assert config is None - @patch("litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper") + @patch("litellm.llms.github_copilot.responses.transformation.cached_get_model_info_helper") def test_returns_none_when_mode_is_unset_and_no_endpoints(self, mock_get_info): """Entry without ``mode`` and without ``supported_endpoints`` returns None (conservative default).""" @@ -463,7 +463,7 @@ class TestGithubCopilotResponsesAPIRouting: ) assert isinstance(config, GithubCopilotResponsesAPIConfig) - @patch("litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper") + @patch("litellm.llms.github_copilot.responses.transformation.cached_get_model_info_helper") def test_returns_none_when_get_model_info_raises(self, mock_get_info): """Catalog lookup failure (model not registered) returns None (conservative default; bridge handles unknown models safely).""" @@ -474,7 +474,7 @@ class TestGithubCopilotResponsesAPIRouting: ) assert config is None - @patch("litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper") + @patch("litellm.llms.github_copilot.responses.transformation.cached_get_model_info_helper") def test_user_override_via_register_model(self, mock_get_info): """User-supplied per-deployment ``model_info`` flows through ``litellm.register_model`` (called by the router) into the merged @@ -488,7 +488,7 @@ class TestGithubCopilotResponsesAPIRouting: ) assert isinstance(config, GithubCopilotResponsesAPIConfig) - @patch("litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper") + @patch("litellm.llms.github_copilot.responses.transformation.cached_get_model_info_helper") def test_realistic_chat_only_entry_returns_none(self, mock_get_info): """Realistic ``model_prices_and_context_window.json`` shape for a chat-only Copilot model (e.g. github_copilot/gemini-3.1-pro-preview) @@ -512,7 +512,7 @@ class TestGithubCopilotResponsesAPIRouting: ) assert config is None - @patch("litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper") + @patch("litellm.llms.github_copilot.responses.transformation.cached_get_model_info_helper") def test_realistic_responses_only_entry_returns_config(self, mock_get_info): """Realistic catalog entry for a Responses-only Copilot model (e.g. github_copilot/gpt-5.5) returns the native config.""" diff --git a/tests/unit/llms/test_file_content_block.py b/tests/unit/llms/test_file_content_block.py index 5552c1a4d68..48e1d92b6cc 100644 --- a/tests/unit/llms/test_file_content_block.py +++ b/tests/unit/llms/test_file_content_block.py @@ -302,12 +302,12 @@ def test_bedrock_process_file_message_malformed_raises_bad_request(): """_process_file_message should raise BadRequestError (not KeyError) when the file object is missing the 'file' sub-field.""" with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"): - BedrockConverseMessagesProcessor._process_file_message(MALFORMED_FILE_OBJECT) + BedrockConverseMessagesProcessor.process_file_message(MALFORMED_FILE_OBJECT) def test_bedrock_process_file_message_explicit_null_file_field_raises_bad_request(): with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"): - BedrockConverseMessagesProcessor._process_file_message(EXPLICIT_NULL_FILE_OBJECT) + BedrockConverseMessagesProcessor.process_file_message(EXPLICIT_NULL_FILE_OBJECT) def test_bedrock_async_process_file_message_malformed_raises_bad_request(): diff --git a/tests/unit/llms/test_file_search_responses.py b/tests/unit/llms/test_file_search_responses.py index 90c60fd20e8..9d2a5dbaf59 100644 --- a/tests/unit/llms/test_file_search_responses.py +++ b/tests/unit/llms/test_file_search_responses.py @@ -330,7 +330,7 @@ class TestManagedFilesVectorStoreAccess: def _make_hook(self): """Return a ManagedFiles instance with prisma_client mocked.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles as ManagedFiles, + PROXY_LiteLLMManagedFiles as ManagedFiles, ) hook = ManagedFiles.__new__(ManagedFiles) @@ -474,7 +474,7 @@ class TestManagedFilesVectorStoreAccess: async def test_F6_non_responses_call_type_skipped(self): """Access check only runs for aresponses/responses call types.""" from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles as ManagedFiles, + PROXY_LiteLLMManagedFiles as ManagedFiles, ) from litellm.proxy._types import CallTypes @@ -502,7 +502,7 @@ class TestManagedFilesVectorStoreAccess: class TestGetVectorStoreIdsFromFileSearchTools: def _make_hook(self): from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles as ManagedFiles, + PROXY_LiteLLMManagedFiles as ManagedFiles, ) return ManagedFiles.__new__(ManagedFiles) diff --git a/tests/unit/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py b/tests/unit/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py index 208cba519f3..6c5c3f8c05f 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py +++ b/tests/unit/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py @@ -14,8 +14,8 @@ import pytest import litellm from litellm.litellm_core_utils.prompt_templates.factory import ( THOUGHT_SIGNATURE_SEPARATOR, - _encode_tool_call_id_with_signature, - _get_thought_signature_from_tool, + encode_tool_call_id_with_signature, + get_thought_signature_from_tool, convert_to_gemini_tool_call_invoke, ) from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( @@ -30,7 +30,7 @@ def test_encode_decode_tool_call_id_with_signature(): test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" # Test encoding - encoded_id = _encode_tool_call_id_with_signature(base_id, test_signature) + encoded_id = encode_tool_call_id_with_signature(base_id, test_signature) assert THOUGHT_SIGNATURE_SEPARATOR in encoded_id assert encoded_id.startswith(base_id) @@ -44,7 +44,7 @@ def test_encode_decode_tool_call_id_with_signature(): }, } - extracted_signature = _get_thought_signature_from_tool(tool) + extracted_signature = get_thought_signature_from_tool(tool) assert extracted_signature == test_signature # Verify base ID is preserved @@ -57,13 +57,13 @@ def test_encode_tool_call_id_without_signature(): base_id = "call_abc123def456" # Encode without signature - encoded_id = _encode_tool_call_id_with_signature(base_id, None) + encoded_id = encode_tool_call_id_with_signature(base_id, None) assert encoded_id == base_id assert THOUGHT_SIGNATURE_SEPARATOR not in encoded_id # Decode ID without signature using factory function tool_obj = {"id": base_id, "type": "function"} - decoded_signature = _get_thought_signature_from_tool(tool_obj) + decoded_signature = get_thought_signature_from_tool(tool_obj) assert decoded_signature is None @@ -103,7 +103,7 @@ def test_tool_call_id_includes_signature_in_response(enable_preview_features): assert THOUGHT_SIGNATURE_SEPARATOR in tool_call_id # Verify we can decode it using the factory function tool_obj = {"id": tool_call_id, "type": "function"} - decoded_sig = _get_thought_signature_from_tool(tool_obj) + decoded_sig = get_thought_signature_from_tool(tool_obj) assert decoded_sig == test_signature @@ -122,7 +122,7 @@ def test_get_thought_signature_backward_compatibility(): "provider_specific_fields": {"thought_signature": test_signature}, } - extracted_signature = _get_thought_signature_from_tool(tool) + extracted_signature = get_thought_signature_from_tool(tool) assert extracted_signature == test_signature @@ -131,7 +131,7 @@ def test_get_thought_signature_prioritizes_provider_fields(): signature_in_fields = "signature_from_fields" signature_in_id = "signature_from_id" - encoded_id = _encode_tool_call_id_with_signature("call_abc123", signature_in_id) + encoded_id = encode_tool_call_id_with_signature("call_abc123", signature_in_id) tool = { "id": encoded_id, @@ -143,7 +143,7 @@ def test_get_thought_signature_prioritizes_provider_fields(): "provider_specific_fields": {"thought_signature": signature_in_fields}, } - extracted_signature = _get_thought_signature_from_tool(tool) + extracted_signature = get_thought_signature_from_tool(tool) # Should prioritize provider_specific_fields assert extracted_signature == signature_in_fields @@ -154,7 +154,7 @@ def test_convert_to_gemini_with_embedded_signature(): # Create tool call ID with embedded signature (as OpenAI client would send) base_id = "call_abc123" - encoded_id = _encode_tool_call_id_with_signature(base_id, test_signature) + encoded_id = encode_tool_call_id_with_signature(base_id, test_signature) # Assistant message as sent by OpenAI client (no provider_specific_fields) assistant_message = { @@ -278,10 +278,10 @@ def test_parallel_tool_calls_with_signatures(enable_preview_features): # When preview features enabled, first tool call has signature in ID assert THOUGHT_SIGNATURE_SEPARATOR in tools[0]["id"] - sig1 = _get_thought_signature_from_tool({"id": tools[0]["id"], "type": "function"}) + sig1 = get_thought_signature_from_tool({"id": tools[0]["id"], "type": "function"}) assert sig1 == signature1 # Second tool call has no signature in ID (regardless of flag) assert THOUGHT_SIGNATURE_SEPARATOR not in tools[1]["id"] - sig2 = _get_thought_signature_from_tool({"id": tools[1]["id"], "type": "function"}) + sig2 = get_thought_signature_from_tool({"id": tools[1]["id"], "type": "function"}) assert sig2 is None diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 940bd71fdb4..114655e6277 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -822,12 +822,12 @@ def _parallel_tool_calls_signed_via_id(*signatures): OpenAI-format client echoes back on the next turn. """ from litellm.litellm_core_utils.prompt_templates.factory import ( - _encode_tool_call_id_with_signature, + encode_tool_call_id_with_signature, ) return [ { - "id": _encode_tool_call_id_with_signature(f"call_{idx}", signature), + "id": encode_tool_call_id_with_signature(f"call_{idx}", signature), "type": "function", "function": {"name": f"tool_{idx}", "arguments": '{"location": "Paris"}'}, "index": idx, diff --git a/tests/unit/llms/vertex_ai/test_vertex.py b/tests/unit/llms/vertex_ai/test_vertex.py index ab8bf123ab2..fe986371b73 100644 --- a/tests/unit/llms/vertex_ai/test_vertex.py +++ b/tests/unit/llms/vertex_ai/test_vertex.py @@ -1271,43 +1271,41 @@ def test_process_gemini_media(): def test_get_image_mime_type_from_url(): - """Test the _get_image_mime_type_from_url function for different image URLs""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _get_image_mime_type_from_url, - ) + """Test MIME type inference for remote media URLs""" + from litellm.litellm_core_utils.prompt_templates.common_utils import get_image_mime_type_from_url # Test JPEG images assert ( - _get_image_mime_type_from_url("https://example.com/image.jpg") == "image/jpeg" + get_image_mime_type_from_url("https://example.com/image.jpg") == "image/jpeg" ) assert ( - _get_image_mime_type_from_url("https://example.com/image.jpeg") == "image/jpeg" + get_image_mime_type_from_url("https://example.com/image.jpeg") == "image/jpeg" ) assert ( - _get_image_mime_type_from_url("https://example.com/IMAGE.JPG") == "image/jpeg" + get_image_mime_type_from_url("https://example.com/IMAGE.JPG") == "image/jpeg" ) # Test PNG images - assert _get_image_mime_type_from_url("https://example.com/image.png") == "image/png" - assert _get_image_mime_type_from_url("https://example.com/IMAGE.PNG") == "image/png" + assert get_image_mime_type_from_url("https://example.com/image.png") == "image/png" + assert get_image_mime_type_from_url("https://example.com/IMAGE.PNG") == "image/png" # Test WebP images assert ( - _get_image_mime_type_from_url("https://example.com/image.webp") == "image/webp" + get_image_mime_type_from_url("https://example.com/image.webp") == "image/webp" ) assert ( - _get_image_mime_type_from_url("https://example.com/IMAGE.WEBP") == "image/webp" + get_image_mime_type_from_url("https://example.com/IMAGE.WEBP") == "image/webp" ) # Test audio formats - assert _get_image_mime_type_from_url("https://example.com/audio.ogg") == "audio/ogg" - assert _get_image_mime_type_from_url("https://example.com/track.OGG") == "audio/ogg" + assert get_image_mime_type_from_url("https://example.com/audio.ogg") == "audio/ogg" + assert get_image_mime_type_from_url("https://example.com/track.OGG") == "audio/ogg" # Test unsupported formats - assert _get_image_mime_type_from_url("https://example.com/image.gif") is None - assert _get_image_mime_type_from_url("https://example.com/image.bmp") is None - assert _get_image_mime_type_from_url("https://example.com/image") is None - assert _get_image_mime_type_from_url("invalid_url") is None + assert get_image_mime_type_from_url("https://example.com/image.gif") is None + assert get_image_mime_type_from_url("https://example.com/image.bmp") is None + assert get_image_mime_type_from_url("https://example.com/image") is None + assert get_image_mime_type_from_url("invalid_url") is None @pytest.mark.parametrize( diff --git a/tests/unit/llms/xai/responses/test_xai_responses_transformation.py b/tests/unit/llms/xai/responses/test_xai_responses_transformation.py index 34ad4b9075d..188ca453364 100644 --- a/tests/unit/llms/xai/responses/test_xai_responses_transformation.py +++ b/tests/unit/llms/xai/responses/test_xai_responses_transformation.py @@ -426,7 +426,7 @@ class TestXAIResponsesWebSearchBilling: def test_bridged_usage_keeps_tool_details_for_billing(self): response = self._transform(include_web_search=True) - bridged = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response.usage) + bridged = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(response.usage) assert isinstance(bridged, Usage) assert bridged.prompt_tokens == 100 @@ -448,7 +448,7 @@ class TestXAIResponsesWebSearchBilling: assert isinstance(event.response.usage, ResponseAPIUsage) assert event.response.usage.input_tokens == 100 - bridged = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(event.response.usage) + bridged = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(event.response.usage) assert getattr(bridged, "server_side_tool_usage_details") == self._TOOL_DETAILS @@ -497,7 +497,7 @@ class TestXAIResponsesReportedCost: assert usage.cost == 0.0037756 - chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + chat_usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert cost_per_token(model="grok-4-latest", usage=chat_usage) == (0.0, 0.0037756) def test_streamed_reported_cost_reaches_the_cost_calculator(self): @@ -519,7 +519,7 @@ class TestXAIResponsesReportedCost: ) assert isinstance(event, ResponseCompletedEvent) - chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(event.response.usage) + chat_usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(event.response.usage) assert cost_per_token(model="grok-4-latest", usage=chat_usage) == (0.0, 0.0037756) def test_usage_without_a_reported_cost_is_left_alone(self): diff --git a/tests/unit/passthrough/test_passthrough_main.py b/tests/unit/passthrough/test_passthrough_main.py index 729b03b7df4..084a0fef9db 100644 --- a/tests/unit/passthrough/test_passthrough_main.py +++ b/tests/unit/passthrough/test_passthrough_main.py @@ -831,7 +831,7 @@ def test_llm_passthrough_route_propagates_allm_passthrough_route_to_logging_obj( result.close() assert captured_litellm_params.get("allm_passthrough_route") is True - assert LitellmLogging._is_sync_litellm_request(captured_litellm_params) is False + assert LitellmLogging.is_sync_litellm_request(captured_litellm_params) is False FOUNDRY_BASE = "https://my-resource.services.ai.azure.com" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_chat_completions.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_chat_completions.py index 79619eefd7f..32cd883c1f7 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_chat_completions.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_chat_completions.py @@ -40,12 +40,12 @@ async def test_acompletion_mcp_auto_exec(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", fake_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", fake_execute, ) monkeypatch.setattr( @@ -103,12 +103,12 @@ async def test_acompletion_mcp_respects_manual_approval(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", fake_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", fake_execute, ) monkeypatch.setattr( @@ -191,12 +191,12 @@ async def test_completion_mcp_with_streaming_no_timeout_error(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", fake_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", fake_execute, ) monkeypatch.setattr( @@ -508,12 +508,12 @@ async def test_mcp_metadata_in_streaming_final_chunk(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", fake_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", fake_execute, ) monkeypatch.setattr( @@ -863,12 +863,12 @@ async def test_mcp_streaming_metadata_ordering(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", fake_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", fake_execute, ) monkeypatch.setattr( diff --git a/tests/unit/proxy/batches_endpoints/test_endpoints.py b/tests/unit/proxy/batches_endpoints/test_endpoints.py index 3bf51f02d34..e2ef8e1c789 100644 --- a/tests/unit/proxy/batches_endpoints/test_endpoints.py +++ b/tests/unit/proxy/batches_endpoints/test_endpoints.py @@ -41,7 +41,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest import respx -from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles +from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles import litellm import litellm.proxy.batches_endpoints.endpoints as endpoints @@ -1166,7 +1166,7 @@ async def test_create__uses_acreate_batch_route_type(harness, openai_env_creds): def install_managed_files_hook(harness: Harness) -> AsyncMock: prisma_client = AsyncMock() - managed_files = _PROXY_LiteLLMManagedFiles(MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client) + managed_files = PROXY_LiteLLMManagedFiles(MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client) harness.logging.post_call_success_hook = AsyncMock(side_effect=managed_files.async_post_call_success_hook) harness.router.model_list = [] return prisma_client diff --git a/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py index cae92ce18bb..ad99cebc342 100644 --- a/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -9,7 +9,7 @@ from unittest.mock import AsyncMock, MagicMock import httpx import pytest -from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles +from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles from openai.types.batch_request_counts import BatchRequestCounts from litellm.models.managed_files import LiteLLM_ManagedFileTable @@ -181,7 +181,7 @@ class FakeManagedBatchStore: return self.objects[unified_batch_id].batch() -REAL_HOOK: Final = _PROXY_LiteLLMManagedFiles(internal_usage_cache=MagicMock(), prisma_client=MagicMock()) +REAL_HOOK: Final = PROXY_LiteLLMManagedFiles(internal_usage_cache=MagicMock(), prisma_client=MagicMock()) class RealIdManagedBatchStore(FakeManagedBatchStore): diff --git a/tests/unit/proxy/common_utils/test_check_batch_cost.py b/tests/unit/proxy/common_utils/test_check_batch_cost.py index 4e7effad3bf..d970978704e 100644 --- a/tests/unit/proxy/common_utils/test_check_batch_cost.py +++ b/tests/unit/proxy/common_utils/test_check_batch_cost.py @@ -106,9 +106,7 @@ class TestCheckBatchCost: return MagicMock() @pytest.fixture - def check_batch_cost_instance( - self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router - ): + def check_batch_cost_instance(self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router): from litellm_enterprise.proxy.common_utils.check_batch_cost import ( CheckBatchCost, ) @@ -120,23 +118,15 @@ class TestCheckBatchCost: ) @pytest.mark.asyncio - async def test_cleanup_scoped_to_batch_file_purpose( - self, check_batch_cost_instance, mock_prisma_client - ): + async def test_cleanup_scoped_to_batch_file_purpose(self, check_batch_cost_instance, mock_prisma_client): """_cleanup_stale_managed_objects scopes its update to file_purpose='batch' only.""" - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) # Return empty so the main poll loop exits immediately - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[]) await check_batch_cost_instance.check_batch_cost() - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list - ) + calls = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list stale_call = calls[0] assert stale_call[1]["data"] == {"status": "stale_expired"} where = stale_call[1]["where"] @@ -145,9 +135,7 @@ class TestCheckBatchCost: assert "created_at" in where @pytest.mark.asyncio - async def test_startup_probe_confirms_batch_processed_support( - self, check_batch_cost_instance, mock_prisma_client - ): + async def test_startup_probe_confirms_batch_processed_support(self, check_batch_cost_instance, mock_prisma_client): mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=None) await check_batch_cost_instance.confirm_batch_processed_support() @@ -158,9 +146,7 @@ class TestCheckBatchCost: assert check_batch_cost_instance._has_batch_processed_column is True @pytest.mark.asyncio - async def test_startup_probe_marks_column_absent( - self, check_batch_cost_instance, mock_prisma_client - ): + async def test_startup_probe_marks_column_absent(self, check_batch_cost_instance, mock_prisma_client): mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( side_effect=Exception("column batch_processed does not exist") ) @@ -184,18 +170,12 @@ class TestCheckBatchCost: assert check_batch_cost_instance._has_batch_processed_column is True @pytest.mark.asyncio - async def test_find_many_uses_pagination_and_excludes_stale( - self, check_batch_cost_instance, mock_prisma_client - ): + async def test_find_many_uses_pagination_and_excludes_stale(self, check_batch_cost_instance, mock_prisma_client): """find_many is called with take, order, and all terminal statuses excluded.""" from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[] - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[]) await check_batch_cost_instance.check_batch_cost() @@ -221,9 +201,7 @@ class TestCheckBatchCost: """Falls back to query without batch_processed when primary query raises.""" from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) # First find_many (primary query) raises with a schema error; second (fallback) returns empty mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( side_effect=[Exception("column batch_processed does not exist"), []] @@ -231,9 +209,7 @@ class TestCheckBatchCost: await check_batch_cost_instance.check_batch_cost() - calls = ( - mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args_list - ) + calls = mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args_list assert len(calls) == 2 fallback_where = calls[1][1]["where"] assert "batch_processed" not in fallback_where @@ -244,32 +220,20 @@ class TestCheckBatchCost: assert check_batch_cost_instance.batch_processed_support_confirmed is False @pytest.mark.asyncio - async def test_column_absence_cached_across_cycles( - self, check_batch_cost_instance, mock_prisma_client - ): + async def test_column_absence_cached_across_cycles(self, check_batch_cost_instance, mock_prisma_client): """After column absence is discovered, subsequent cycles skip the primary query entirely.""" from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) # Simulate column already known absent from a previous cycle check_batch_cost_instance._has_batch_processed_column = False - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[]) await check_batch_cost_instance.check_batch_cost() # Only one find_many call — the fallback directly, no primary query attempt - assert ( - mock_prisma_client.db.litellm_managedobjecttable.find_many.call_count == 1 - ) - fallback_where = ( - mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args[1][ - "where" - ] - ) + assert mock_prisma_client.db.litellm_managedobjecttable.find_many.call_count == 1 + fallback_where = mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args[1]["where"] assert "batch_processed" not in fallback_where @pytest.mark.asyncio @@ -283,13 +247,9 @@ class TestCheckBatchCost: """ from unittest.mock import patch - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_job = MagicMock() mock_job.id = "job-fallback-1" @@ -298,22 +258,16 @@ class TestCheckBatchCost: # Simulate column already known absent (e.g. discovered on a previous cycle) check_batch_cost_instance._has_batch_processed_column = False - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) # Build a fake batch response whose status triggers the completion branch mock_response = MagicMock() mock_response.status = "completed" mock_response.output_file_id = "file-output-123" - mock_response.model_dump_json.return_value = ( - '{"id":"batch-1","status":"completed"}' - ) + mock_response.model_dump_json.return_value = '{"id":"batch-1","status":"completed"}' mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) - mock_llm_router.get_deployment_credentials_with_provider = MagicMock( - return_value={"api_key": "sk-test"} - ) + mock_llm_router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) mock_deployment = MagicMock() mock_deployment.litellm_params.custom_llm_provider = "openai" @@ -345,7 +299,7 @@ class TestCheckBatchCost: return_value=mock_file_content, ), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"id": "req-1"}], ), patch( @@ -361,9 +315,7 @@ class TestCheckBatchCost: "litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", return_value=("gpt-4", "openai", None, None), ), - patch( - "litellm.litellm_core_utils.litellm_logging.Logging" - ) as mock_logging_cls, + patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls, ): mock_logging_obj = MagicMock() mock_logging_obj.async_success_handler = AsyncMock() @@ -372,15 +324,11 @@ class TestCheckBatchCost: await check_batch_cost_instance.check_batch_cost() # The update must have been called — this is the core assertion. - assert ( - mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1 - ), "Expected update() to be called exactly once for the completed job" - update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[ - 1 - ]["data"] - assert ( - "batch_processed" not in update_data - ), "update() must NOT include batch_processed when column is absent" + assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1, ( + "Expected update() to be called exactly once for the completed job" + ) + update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[1]["data"] + assert "batch_processed" not in update_data, "update() must NOT include batch_processed when column is absent" assert update_data["status"] == "complete" @pytest.mark.asyncio @@ -450,13 +398,15 @@ class TestCheckBatchCost: return_value=mock_file_content, ) as mock_afile_content, patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"recordId": "req-1"}], ), patch( "litellm.batches.batch_utils.calculate_batch_cost_and_usage", new_callable=AsyncMock, - return_value=_batch_cost_result(0.01, {"prompt_tokens": 10, "completion_tokens": 5}, ["claude-haiku-4-5"]), + return_value=_batch_cost_result( + 0.01, {"prompt_tokens": 10, "completion_tokens": 5}, ["claude-haiku-4-5"] + ), ), patch( "litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", @@ -474,9 +424,9 @@ class TestCheckBatchCost: passed_kwargs = mock_afile_content.await_args[1] snapshot = passed_kwargs.get("_litellm_internal_model_credentials") assert snapshot is not None, "cost poller must pass the trusted credential snapshot" - assert isinstance( - snapshot, MappingProxyType - ), "snapshot must be a MappingProxyType; a plain dict is rejected by get_configured_s3_bucket_name" + assert isinstance(snapshot, MappingProxyType), ( + "snapshot must be a MappingProxyType; a plain dict is rejected by get_configured_s3_bucket_name" + ) assert snapshot["s3_bucket_name"] == "configured-batch-bucket" @pytest.mark.asyncio @@ -553,7 +503,7 @@ class TestCheckBatchCost: return_value=mock_file_content, ), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"recordId": "req-1"}], ), patch( @@ -678,13 +628,9 @@ class TestCheckBatchCost: """ from unittest.mock import patch - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_job = MagicMock() mock_job.id = "job-primary-1" @@ -692,21 +638,15 @@ class TestCheckBatchCost: mock_job.created_by = "user-1" assert check_batch_cost_instance._has_batch_processed_column is True - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) mock_response = MagicMock() mock_response.status = "completed" mock_response.output_file_id = "file-output-123" - mock_response.model_dump_json.return_value = ( - '{"id":"batch-1","status":"completed"}' - ) + mock_response.model_dump_json.return_value = '{"id":"batch-1","status":"completed"}' mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) - mock_llm_router.get_deployment_credentials_with_provider = MagicMock( - return_value={"api_key": "sk-test"} - ) + mock_llm_router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) mock_deployment = MagicMock() mock_deployment.litellm_params.custom_llm_provider = "openai" @@ -738,7 +678,7 @@ class TestCheckBatchCost: return_value=mock_file_content, ), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"id": "req-1"}], ), patch( @@ -754,9 +694,7 @@ class TestCheckBatchCost: "litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", return_value=("gpt-4", "openai", None, None), ), - patch( - "litellm.litellm_core_utils.litellm_logging.Logging" - ) as mock_logging_cls, + patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls, ): mock_logging_obj = MagicMock() mock_logging_obj.async_success_handler = AsyncMock() @@ -764,15 +702,13 @@ class TestCheckBatchCost: await check_batch_cost_instance.check_batch_cost() - assert ( - mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1 - ), "Expected update() to be called exactly once for the completed job" - update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[ - 1 - ]["data"] - assert ( - update_data["batch_processed"] is True - ), "update() must include batch_processed=True when column is present" + assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1, ( + "Expected update() to be called exactly once for the completed job" + ) + update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[1]["data"] + assert update_data["batch_processed"] is True, ( + "update() must include batch_processed=True when column is present" + ) assert update_data["status"] == "complete" @pytest.mark.asyncio @@ -868,7 +804,7 @@ class TestCheckBatchCost: return_value=mock_file_content, ), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"id": "req-1"}], ), patch( @@ -916,22 +852,16 @@ class TestCheckBatchCost: """ from unittest.mock import patch - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_job = MagicMock() mock_job.id = "job-anthropic-1" mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA==" mock_job.created_by = "user-1" - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) mock_response = MagicMock() mock_response.status = "completed" @@ -965,9 +895,9 @@ class TestCheckBatchCost: ): await check_batch_cost_instance.check_batch_cost() - assert ( - mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0 - ), "a failed cost tracking attempt must not mark the job processed" + assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0, ( + "a failed cost tracking attempt must not mark the job processed" + ) @pytest.mark.asyncio @pytest.mark.parametrize("terminal_status", ["failed", "expired", "cancelled"]) @@ -984,13 +914,9 @@ class TestCheckBatchCost: """ import base64 - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_job = MagicMock() mock_job.id = "job-terminal-1" @@ -1000,31 +926,25 @@ class TestCheckBatchCost: mock_job.created_by = "user-1" assert check_batch_cost_instance._has_batch_processed_column is True - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) mock_response = MagicMock() mock_response.status = terminal_status mock_response.output_file_id = None - mock_response.model_dump_json.return_value = ( - f'{{"id":"batch-1","status":"{terminal_status}"}}' - ) + mock_response.model_dump_json.return_value = f'{{"id":"batch-1","status":"{terminal_status}"}}' mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) await check_batch_cost_instance.check_batch_cost() - assert ( - mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1 - ), f"Expected update() to be called exactly once for a {terminal_status} job" - update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[ - 1 - ]["data"] + assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1, ( + f"Expected update() to be called exactly once for a {terminal_status} job" + ) + update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[1]["data"] assert update_data["status"] == terminal_status - assert ( - update_data["batch_processed"] is True - ), "terminal-status update() must set batch_processed=True so polling stops" + assert update_data["batch_processed"] is True, ( + "terminal-status update() must set batch_processed=True so polling stops" + ) @pytest.mark.asyncio @pytest.mark.parametrize("terminal_status", ["failed", "cancelled"]) @@ -1059,13 +979,9 @@ class TestCheckBatchCost: f"litellm_proxy:application/octet-stream;unified_id,u-2;llm_output_file_id,{raw_error_file_id}".encode() ).decode() - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) input_file_row = MagicMock() input_file_row.unified_file_id = unified_input_file_id @@ -1075,9 +991,7 @@ class TestCheckBatchCost: return input_file_row return None - mock_prisma_client.db.litellm_managedfiletable.find_first = AsyncMock( - side_effect=find_managed_file - ) + mock_prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(side_effect=find_managed_file) mock_job = MagicMock() mock_job.id = "job-terminal-mint-1" @@ -1086,9 +1000,7 @@ class TestCheckBatchCost: mock_job.team_id = "team-1" check_batch_cost_instance._has_batch_processed_column = True - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) response = LiteLLMBatch( id="batch-456", @@ -1106,9 +1018,7 @@ class TestCheckBatchCost: mock_hook = MagicMock() mock_hook.get_unified_output_file_id.side_effect = [unified_error_file_id] mock_hook.store_unified_file_id = AsyncMock() - check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = ( - mock_hook - ) + check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = mock_hook await check_batch_cost_instance.check_batch_cost() @@ -1161,13 +1071,9 @@ class TestCheckBatchCost: import base64 from unittest.mock import patch - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_job = MagicMock() mock_job.id = "job-completed-no-output-1" @@ -1177,25 +1083,19 @@ class TestCheckBatchCost: mock_job.created_by = "user-1" assert check_batch_cost_instance._has_batch_processed_column is True - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) mock_response = MagicMock() mock_response.status = completed_status mock_response.output_file_id = None mock_response.error_file_id = "file-error-123" mock_response.request_counts = MagicMock(completed=0, failed=3, total=3) - mock_response.model_dump_json.return_value = ( - f'{{"id":"batch-1","status":"{completed_status}"}}' - ) + mock_response.model_dump_json.return_value = f'{{"id":"batch-1","status":"{completed_status}"}}' mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) # Billing reads credentials off the router; if it is touched we billed a batch # that has no output, which is the behaviour this test guards against. - mock_llm_router.get_deployment_credentials_with_provider = MagicMock( - return_value={"api_key": "sk-test"} - ) + mock_llm_router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) with patch( "litellm.files.main.afile_content", @@ -1203,22 +1103,18 @@ class TestCheckBatchCost: ) as mock_afile_content: await check_batch_cost_instance.check_batch_cost() - assert ( - mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1 - ), "a completed batch with no output file must be marked processed exactly once" - update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[ - 1 - ]["data"] + assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1, ( + "a completed batch with no output file must be marked processed exactly once" + ) + update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[1]["data"] assert update_data["status"] == completed_status - assert ( - update_data["batch_processed"] is True - ), "completed-without-output update() must set batch_processed=True so polling stops" - assert ( - mock_afile_content.await_count == 0 - ), "a batch with no output file must not be billed" - assert ( - mock_llm_router.get_deployment_credentials_with_provider.call_count == 0 - ), "a batch with no output file must not enter the cost-tracking path" + assert update_data["batch_processed"] is True, ( + "completed-without-output update() must set batch_processed=True so polling stops" + ) + assert mock_afile_content.await_count == 0, "a batch with no output file must not be billed" + assert mock_llm_router.get_deployment_credentials_with_provider.call_count == 0, ( + "a batch with no output file must not enter the cost-tracking path" + ) @pytest.mark.asyncio @pytest.mark.parametrize( @@ -1246,13 +1142,9 @@ class TestCheckBatchCost: import base64 from unittest.mock import patch - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_job = MagicMock() mock_job.id = "job-completed-lagging-output-1" @@ -1262,9 +1154,7 @@ class TestCheckBatchCost: mock_job.created_by = "user-1" assert check_batch_cost_instance._has_batch_processed_column is True - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) mock_response = MagicMock() mock_response.status = "completed" @@ -1273,9 +1163,7 @@ class TestCheckBatchCost: mock_response.request_counts = request_counts mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) - mock_llm_router.get_deployment_credentials_with_provider = MagicMock( - return_value={"api_key": "sk-test"} - ) + mock_llm_router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) with patch( "litellm.files.main.afile_content", @@ -1283,12 +1171,10 @@ class TestCheckBatchCost: ) as mock_afile_content: await check_batch_cost_instance.check_batch_cost() - assert ( - mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0 - ), "a completed batch whose output id is still lagging must stay eligible for the next poll" - assert ( - mock_afile_content.await_count == 0 - ), "a batch with no output file must not be billed" + assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0, ( + "a completed batch whose output id is still lagging must stay eligible for the next poll" + ) + assert mock_afile_content.await_count == 0, "a batch with no output file must not be billed" @pytest.mark.asyncio async def test_non_terminal_status_left_unprocessed( @@ -1299,9 +1185,7 @@ class TestCheckBatchCost: """ from unittest.mock import patch - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_job = MagicMock() @@ -1309,9 +1193,7 @@ class TestCheckBatchCost: mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA==" mock_job.created_by = "user-1" - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) mock_response = MagicMock() mock_response.status = "in_progress" @@ -1337,9 +1219,9 @@ class TestCheckBatchCost: ): await check_batch_cost_instance.check_batch_cost() - assert ( - mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0 - ), "a non-terminal batch must not be written back (would stop polling prematurely)" + assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0, ( + "a non-terminal batch must not be written back (would stop polling prematurely)" + ) @pytest.mark.asyncio @pytest.mark.parametrize("terminal_status", ["expired", "cancelled", "failed"]) @@ -1356,13 +1238,9 @@ class TestCheckBatchCost: """ from unittest.mock import patch - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_job = MagicMock() mock_job.id = "job-terminal-with-output-1" @@ -1370,21 +1248,15 @@ class TestCheckBatchCost: mock_job.created_by = "user-1" assert check_batch_cost_instance._has_batch_processed_column is True - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) mock_response = MagicMock() mock_response.status = terminal_status mock_response.output_file_id = "file-output-123" - mock_response.model_dump_json.return_value = ( - f'{{"id":"batch-1","status":"{terminal_status}"}}' - ) + mock_response.model_dump_json.return_value = f'{{"id":"batch-1","status":"{terminal_status}"}}' mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) - mock_llm_router.get_deployment_credentials_with_provider = MagicMock( - return_value={"api_key": "sk-test"} - ) + mock_llm_router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) mock_deployment = MagicMock() mock_deployment.litellm_params.custom_llm_provider = "openai" @@ -1416,7 +1288,7 @@ class TestCheckBatchCost: return_value=mock_file_content, ) as mock_afile_content, patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"id": "req-1"}], ), patch( @@ -1432,9 +1304,7 @@ class TestCheckBatchCost: "litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", return_value=("gpt-4", "openai", None, None), ), - patch( - "litellm.litellm_core_utils.litellm_logging.Logging" - ) as mock_logging_cls, + patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls, ): mock_logging_obj = MagicMock() mock_logging_obj.async_success_handler = AsyncMock() @@ -1442,20 +1312,16 @@ class TestCheckBatchCost: await check_batch_cost_instance.check_batch_cost() - assert ( - mock_afile_content.await_count == 1 - ), f"{terminal_status} batch with an output file must fetch results and be billed" - mock_logging_obj.async_success_handler.assert_awaited_once() - assert ( - mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1 + assert mock_afile_content.await_count == 1, ( + f"{terminal_status} batch with an output file must fetch results and be billed" ) - update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[ - 1 - ]["data"] + mock_logging_obj.async_success_handler.assert_awaited_once() + assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1 + update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[1]["data"] assert update_data["batch_processed"] is True - assert ( - update_data["status"] == terminal_status - ), f"billed {terminal_status} batch must keep its real terminal status in the DB" + assert update_data["status"] == terminal_status, ( + f"billed {terminal_status} batch must keep its real terminal status in the DB" + ) @pytest.mark.asyncio async def test_error_file_failures_add_to_failed_request_count( @@ -1580,13 +1446,9 @@ class TestCheckBatchCost: from litellm.exceptions import NotFoundError - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_job = MagicMock() mock_job.id = "job-output-gone-1" @@ -1596,23 +1458,17 @@ class TestCheckBatchCost: mock_job.created_by = "user-1" assert check_batch_cost_instance._has_batch_processed_column is True - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) missing_output_file_id = "gs://batch-out/job-1/predictions.jsonl" mock_response = MagicMock() mock_response.status = "failed" mock_response.output_file_id = missing_output_file_id mock_response.error_file_id = None - mock_response.model_dump_json.return_value = ( - '{"id":"batch-1","status":"failed"}' - ) + mock_response.model_dump_json.return_value = '{"id":"batch-1","status":"failed"}' mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) - mock_llm_router.get_deployment_credentials_with_provider = MagicMock( - return_value={"api_key": "sk-test"} - ) + mock_llm_router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) with ( patch( @@ -1633,12 +1489,10 @@ class TestCheckBatchCost: assert mock_afile_content.await_count == 1 mock_calculate.assert_not_awaited() - assert ( - mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1 - ), "a terminal batch with a 404ing output file must be retired, not retried forever" - update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[ - 1 - ]["data"] + assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1, ( + "a terminal batch with a 404ing output file must be retired, not retried forever" + ) + update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[1]["data"] assert update_data["status"] == "failed" assert update_data["batch_processed"] is True @@ -1651,13 +1505,9 @@ class TestCheckBatchCost: Without this, GET /batches/{id} returns a raw file ID that cannot be routed through the proxy, causing API_KEY errors when clients call GET /files/{id}/content. """ - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=1 - ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) mock_job = MagicMock() mock_job.id = "job-raw-file-1" @@ -1666,9 +1516,7 @@ class TestCheckBatchCost: mock_job.team_id = None check_batch_cost_instance._has_batch_processed_column = True - mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[mock_job] - ) + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) raw_output_file_id = "file-batch-output-abc123" raw_error_file_id = "file-batch-error-xyz456" @@ -1679,14 +1527,10 @@ class TestCheckBatchCost: mock_response.status = "completed" mock_response.output_file_id = raw_output_file_id mock_response.error_file_id = raw_error_file_id - mock_response.model_dump_json.return_value = ( - '{"id":"batch-1","status":"completed"}' - ) + mock_response.model_dump_json.return_value = '{"id":"batch-1","status":"completed"}' mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) - mock_llm_router.get_deployment_credentials_with_provider = MagicMock( - return_value={"api_key": "sk-test"} - ) + mock_llm_router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) mock_deployment = MagicMock() mock_deployment.litellm_params.custom_llm_provider = "azure" @@ -1701,9 +1545,7 @@ class TestCheckBatchCost: fake_managed_error_id, ] mock_hook.store_unified_file_id = AsyncMock() - check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = ( - mock_hook - ) + check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = mock_hook mock_file_content = MagicMock() mock_file_content.content = b'{"id":"req-1"}' @@ -1730,7 +1572,7 @@ class TestCheckBatchCost: return_value=mock_file_content, ), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"id": "req-1"}], ), patch( @@ -1746,9 +1588,7 @@ class TestCheckBatchCost: "litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", return_value=("gpt-5-mini", "azure", None, None), ), - patch( - "litellm.litellm_core_utils.litellm_logging.Logging" - ) as mock_logging_cls, + patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls, ): mock_logging_obj = MagicMock() mock_logging_obj.async_success_handler = AsyncMock() @@ -1799,9 +1639,7 @@ class TestUnmanagedVertexRouting: def _job(self, file_object=None): job = MagicMock() job.unified_object_id = "8823717160934178816" - job.file_object = ( - file_object if file_object is not None else _unmanaged_vertex_file_object() - ) + job.file_object = file_object if file_object is not None else _unmanaged_vertex_file_object() return job def test_flag_off_skips_unmanaged_id_unchanged(self): @@ -1839,9 +1677,7 @@ class TestUnmanagedVertexRouting: assert result == ("deploy-1", "8823717160934178816") # bare model name (trailing GCS segment), not the full publishers/.. path - router.resolve_model_name_from_model_id.assert_called_once_with( - "gemini-2.5-flash" - ) + router.resolve_model_name_from_model_id.assert_called_once_with("gemini-2.5-flash") router.get_model_ids.assert_called_once_with(model_name="gemini-2.5-flash") def test_flag_on_routes_fine_tuned_endpoint_to_vertex_deployment(self): @@ -1890,9 +1726,7 @@ class TestUnmanagedVertexRouting: result = instance._resolve_job_routing(self._job(), prom) assert result is None - prom.record_check_batch_cost_error.assert_called_once_with( - "unmanaged_no_matching_deployment" - ) + prom.record_check_batch_cost_error.assert_called_once_with("unmanaged_no_matching_deployment") def test_flag_on_uses_later_vertex_deployment_with_matching_suffix(self): router = MagicMock() @@ -1940,9 +1774,7 @@ class TestUnmanagedVertexRouting: result = instance._resolve_job_routing(self._job(), prom) assert result is None - prom.record_check_batch_cost_error.assert_called_once_with( - "unmanaged_no_matching_deployment" - ) + prom.record_check_batch_cost_error.assert_called_once_with("unmanaged_no_matching_deployment") def test_flag_on_non_gcs_input_is_not_unmanaged_vertex(self): """Flag on, but input_file_id is not a gs:// publishers path: treat as unroutable, @@ -1950,9 +1782,7 @@ class TestUnmanagedVertexRouting: router = MagicMock() instance = self._instance(track_unmanaged=True, router=router) prom = MagicMock() - job = self._job( - file_object=_unmanaged_vertex_file_object(input_file_id="file-abc-123") - ) + job = self._job(file_object=_unmanaged_vertex_file_object(input_file_id="file-abc-123")) with patch(_IS_B64, return_value=False): result = instance._resolve_job_routing(job, prom) @@ -1976,9 +1806,7 @@ class TestUnmanagedVertexRouting: mock_response.error_file_id = None mock_response.completed_at = None mock_response.created_at = None - mock_response.model_dump_json.return_value = ( - '{"id":"8823717160934178816","status":"completed"}' - ) + mock_response.model_dump_json.return_value = '{"id":"8823717160934178816","status":"completed"}' router.aretrieve_batch = AsyncMock(return_value=mock_response) router.get_deployment_credentials_with_provider = MagicMock( return_value={"vertex_project": "p", "vertex_location": "us-central1"} @@ -2000,9 +1828,7 @@ class TestUnmanagedVertexRouting: prisma.db.litellm_managedobjecttable = MagicMock() prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) prisma.db.litellm_managedobjecttable.update = AsyncMock() - prisma.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[self._job()] - ) + prisma.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[self._job()]) prisma.db.litellm_usertable = MagicMock() prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None) @@ -2017,7 +1843,7 @@ class TestUnmanagedVertexRouting: return_value=mock_file_content, ), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"id": "req-1"}], ), patch( @@ -2033,9 +1859,7 @@ class TestUnmanagedVertexRouting: "litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", return_value=("gemini-2.5-flash", "vertex_ai", None, None), ), - patch( - "litellm.litellm_core_utils.litellm_logging.Logging" - ) as mock_logging_cls, + patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls, ): mock_logging_obj = MagicMock() mock_logging_obj.async_success_handler = AsyncMock() @@ -2076,9 +1900,7 @@ class TestUnmanagedBedrockRouting: def _job(self, file_object=None): job = MagicMock() job.unified_object_id = self._ARN - job.file_object = ( - file_object if file_object is not None else _unmanaged_bedrock_file_object() - ) + job.file_object = file_object if file_object is not None else _unmanaged_bedrock_file_object() return job def _bedrock_deployment(self): @@ -2133,9 +1955,7 @@ class TestUnmanagedBedrockRouting: result = instance._resolve_job_routing(self._job(), prom) assert result is None - prom.record_check_batch_cost_error.assert_called_once_with( - "unmanaged_no_matching_deployment" - ) + prom.record_check_batch_cost_error.assert_called_once_with("unmanaged_no_matching_deployment") def test_flag_on_matches_deployment_despite_colon_dash_mismatch(self): """The S3 object key has ':' replaced with '-' (e.g. 'v1-0'), but the configured @@ -2173,9 +1993,7 @@ class TestUnmanagedBedrockRouting: result = instance._resolve_job_routing(self._job(), prom) assert result is None - prom.record_check_batch_cost_error.assert_called_once_with( - "unmanaged_no_matching_deployment" - ) + prom.record_check_batch_cost_error.assert_called_once_with("unmanaged_no_matching_deployment") def test_flag_on_non_s3_input_is_not_unmanaged_bedrock(self): """Flag on, but input_file_id is not a litellm-bedrock-files- s3:// key: treat as @@ -2183,9 +2001,7 @@ class TestUnmanagedBedrockRouting: router = MagicMock() instance = self._instance(track_unmanaged=True, router=router) prom = MagicMock() - job = self._job( - file_object=_unmanaged_bedrock_file_object(input_file_id="file-abc-123") - ) + job = self._job(file_object=_unmanaged_bedrock_file_object(input_file_id="file-abc-123")) with patch(_IS_B64, return_value=False): result = instance._resolve_job_routing(job, prom) @@ -2208,13 +2024,9 @@ class TestUnmanagedBedrockRouting: mock_response.error_file_id = None mock_response.completed_at = None mock_response.created_at = None - mock_response.model_dump_json.return_value = ( - f'{{"id":"{self._ARN}","status":"completed"}}' - ) + mock_response.model_dump_json.return_value = f'{{"id":"{self._ARN}","status":"completed"}}' router.aretrieve_batch = AsyncMock(return_value=mock_response) - router.get_deployment_credentials_with_provider = MagicMock( - return_value={"aws_region_name": "us-east-1"} - ) + router.get_deployment_credentials_with_provider = MagicMock(return_value={"aws_region_name": "us-east-1"}) deployment = self._bedrock_deployment() deployment.model_name = "claude-sonnet-4" @@ -2230,9 +2042,7 @@ class TestUnmanagedBedrockRouting: prisma.db.litellm_managedobjecttable = MagicMock() prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) prisma.db.litellm_managedobjecttable.update = AsyncMock() - prisma.db.litellm_managedobjecttable.find_many = AsyncMock( - return_value=[self._job()] - ) + prisma.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[self._job()]) prisma.db.litellm_usertable = MagicMock() prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None) @@ -2247,7 +2057,7 @@ class TestUnmanagedBedrockRouting: return_value=mock_file_content, ), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"id": "req-1"}], ), patch( @@ -2263,9 +2073,7 @@ class TestUnmanagedBedrockRouting: "litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", return_value=("claude-sonnet-4", "bedrock", None, None), ), - patch( - "litellm.litellm_core_utils.litellm_logging.Logging" - ) as mock_logging_cls, + patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls, ): mock_logging_obj = MagicMock() mock_logging_obj.async_success_handler = AsyncMock() @@ -2385,13 +2193,11 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: ) from litellm.types.utils import LiteLLMBatch from enterprise.litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) router = MagicMock() - router.get_deployment_credentials_with_provider = MagicMock( - return_value={"api_key": "sk-test"} - ) + router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) deployment = MagicMock() deployment.litellm_params.custom_llm_provider = "azure" deployment.litellm_params.model = "azure/gpt-5.5" @@ -2400,8 +2206,8 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: router.get_deployment = MagicMock(return_value=deployment) hook = MagicMock() - hook.get_unified_output_file_id = ( - lambda output_file_id, model_id, model_name: _PROXY_LiteLLMManagedFiles.get_unified_output_file_id( + hook.get_unified_output_file_id = lambda output_file_id, model_id, model_name: ( + PROXY_LiteLLMManagedFiles.get_unified_output_file_id( None, output_file_id=output_file_id, model_id=model_id, model_name=model_name ) ) @@ -2439,7 +2245,7 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: return_value=file_content, ), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"id": "req-1"}], ), patch( @@ -2470,9 +2276,7 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: get_models_from_unified_file_id, ) - output_file_id = await self._run( - self._job(self._managed_input_file_id(self._PUBLIC_MODEL_GROUP)) - ) + output_file_id = await self._run(self._job(self._managed_input_file_id(self._PUBLIC_MODEL_GROUP))) decoded = _is_base64_encoded_unified_file_id(output_file_id) assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP] @@ -2486,9 +2290,7 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: _extract_models_from_managed_resource_id, ) - output_file_id = await self._run( - self._job(self._managed_input_file_id(self._PUBLIC_MODEL_GROUP)) - ) + output_file_id = await self._run(self._job(self._managed_input_file_id(self._PUBLIC_MODEL_GROUP))) models = _extract_models_from_managed_resource_id(output_file_id, "file_id", None) assert models == [self._PUBLIC_MODEL_GROUP] @@ -2496,9 +2298,7 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: await can_key_call_model( model=models[0], llm_model_list=None, - valid_token=UserAPIKeyAuth( - api_key="sk-test", models=[self._PUBLIC_MODEL_GROUP] - ), + valid_token=UserAPIKeyAuth(api_key="sk-test", models=[self._PUBLIC_MODEL_GROUP]), llm_router=None, ) is True @@ -2515,6 +2315,8 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: decoded = _is_base64_encoded_unified_file_id(output_file_id) assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP] + + class TestBatchCostAttribution: """CheckBatchCost rebuilds the creator's spend metadata from the managed-object row so the batch-cost log is attributed like a non-batch request.""" @@ -2611,9 +2413,7 @@ class TestBatchCostAttribution: """An alias lookup failure must not lose the spend row; the key hash and team still attribute it.""" instance = self._instance() - instance.prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( - side_effect=Exception("db down") - ) + instance.prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(side_effect=Exception("db down")) metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") @@ -2624,9 +2424,7 @@ class TestBatchCostAttribution: async def test_cli_session_batch_keeps_its_alias_without_a_key_row(self): instance = self._instance(key_row=None) - metadata = await instance._build_creator_attribution_metadata( - self._job(api_key="cli-session-alice"), "batch-1" - ) + metadata = await instance._build_creator_attribution_metadata(self._job(api_key="cli-session-alice"), "batch-1") assert metadata["user_api_key"] == "cli-session-alice" assert metadata["user_api_key_alias"] == "cli-session-alice" @@ -2699,9 +2497,7 @@ class TestBatchCostAttribution: team_row=SimpleNamespace(team_alias="Team Alpha", organization_id="org-team"), ) - metadata = await instance._build_creator_attribution_metadata( - self._job(org_id="org-at-creation"), "batch-1" - ) + metadata = await instance._build_creator_attribution_metadata(self._job(org_id="org-at-creation"), "batch-1") assert metadata["user_api_key_org_id"] == "org-at-creation" @@ -2744,9 +2540,7 @@ class TestBatchCostAttribution: instance = self._instance( team_row=SimpleNamespace(team_alias="Team Alpha", organization_id="org-team"), ) - instance.prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( - side_effect=Exception("db down") - ) + instance.prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(side_effect=Exception("db down")) metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") @@ -2786,9 +2580,7 @@ class TestBatchCostAttribution: key_row=SimpleNamespace(key_alias="prod-key"), user_row=SimpleNamespace(user_email="alice@example.com", user_alias=None), ) - metadata = await instance._build_creator_attribution_metadata( - self._job(api_key=token_hash), "batch-1" - ) + metadata = await instance._build_creator_attribution_metadata(self._job(api_key=token_hash), "batch-1") assert metadata["user_api_key"] == token_hash assert metadata["user_api_key_hash"] == token_hash @@ -2851,9 +2643,7 @@ class TestPollPageStarvation: async def test_unified_id_without_model_id_is_retired(self): """A unified id that decodes but carries no model_id is unroutable no matter what the config says, so it must leave the poll page instead of being retried forever.""" - prisma = self._prisma( - [self._job("job-no-model", self._encode("litellm_proxy;llm_batch_id:poison-no-model"))] - ) + prisma = self._prisma([self._job("job-no-model", self._encode("litellm_proxy;llm_batch_id:poison-no-model"))]) llm_router = MagicMock() llm_router.aretrieve_batch = AsyncMock() @@ -2891,9 +2681,7 @@ class TestPollPageStarvation: await self._instance(prisma, llm_router).check_batch_cost() prisma.db.litellm_managedobjecttable.update.assert_awaited_once() - assert prisma.db.litellm_managedobjecttable.update.call_args[1]["data"] == { - "batch_processed": True - } + assert prisma.db.litellm_managedobjecttable.update.call_args[1]["data"] == {"batch_processed": True} @pytest.mark.asyncio async def test_provider_404_with_deployment_gone_keeps_job(self): @@ -2946,17 +2734,13 @@ class TestPollPageStarvation: async def test_retirement_falls_back_to_status_without_batch_processed_column(self): """Older schemas have no batch_processed column, so the only way to stop selecting the row is the status filter the poll query already applies.""" - prisma = self._prisma( - [self._job("job-legacy", self._encode("litellm_proxy;llm_batch_id:poison-no-model"))] - ) + prisma = self._prisma([self._job("job-legacy", self._encode("litellm_proxy;llm_batch_id:poison-no-model"))]) instance = self._instance(prisma, MagicMock()) instance._has_batch_processed_column = False await instance.check_batch_cost() - assert prisma.db.litellm_managedobjecttable.update.call_args[1]["data"] == { - "status": "stale_expired" - } + assert prisma.db.litellm_managedobjecttable.update.call_args[1]["data"] == {"status": "stale_expired"} @pytest.mark.asyncio async def test_stale_cleanup_gives_up_on_never_costed_completed_rows(self): @@ -3011,14 +2795,11 @@ class TestPollPageStarvation: await self._instance(prisma, llm_router).check_batch_cost() - retired = [ - call[1]["where"]["id"] - for call in prisma.db.litellm_managedobjecttable.update.call_args_list - ] + retired = [call[1]["where"]["id"] for call in prisma.db.litellm_managedobjecttable.update.call_args_list] assert retired == ["job-no-model", "job-gone"] - assert ( - llm_router.aretrieve_batch.await_args_list[-1][1]["batch_id"] == "batch_live" - ), "the newer healthy batch must still be polled in the same cycle" + assert llm_router.aretrieve_batch.await_args_list[-1][1]["batch_id"] == "batch_live", ( + "the newer healthy batch must still be polled in the same cycle" + ) @pytest.mark.asyncio async def test_404_that_does_not_name_the_batch_keeps_job_for_retry(self): @@ -3047,6 +2828,7 @@ class TestPollPageStarvation: prisma.db.litellm_managedobjecttable.update.assert_not_awaited() + class _FakeManagedObjectRow: """One managed batch row the provider has finished but nothing has costed yet.""" @@ -3063,8 +2845,12 @@ class _FakeManagedObjectRow: self.request_tags = None self.created_at = 1700000000 self.file_object = json.dumps( - {"id": "batch-456", "status": "in_progress", "input_file_id": "file-input-1", - "output_file_id": _CLAIM_OUTPUT_FILE_ID} + { + "id": "batch-456", + "status": "in_progress", + "input_file_id": "file-input-1", + "output_file_id": _CLAIM_OUTPUT_FILE_ID, + } ) @@ -3164,9 +2950,7 @@ class TestMultiPodBatchCostClaim: router = MagicMock() router.aretrieve_batch = AsyncMock(return_value=response) - router.get_deployment_credentials_with_provider = MagicMock( - return_value={"api_key": "sk-test"} - ) + router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) router.get_deployment = MagicMock(return_value=deployment) return router @@ -3210,7 +2994,7 @@ class TestMultiPodBatchCostClaim: ), patch("litellm.files.main.afile_content", new=AsyncMock(side_effect=_afile_content)), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary", + "litellm.batches.batch_utils.get_file_content_as_dictionary", return_value=[{"id": "req-1"}], ), patch( @@ -3237,12 +3021,12 @@ class TestMultiPodBatchCostClaim: @staticmethod async def _run_deletion_guard(prisma, file_id: str) -> None: """Run the real managed-files deletion guard against the row the poller is costing.""" - from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles cache = MagicMock() cache.async_get_cache = AsyncMock(return_value=None) cache.async_set_cache = AsyncMock() - guard = _PROXY_LiteLLMManagedFiles(internal_usage_cache=cache, prisma_client=prisma) + guard = PROXY_LiteLLMManagedFiles(internal_usage_cache=cache, prisma_client=prisma) scheduler = MagicMock() scheduler.get_job.return_value = MagicMock() @@ -3326,9 +3110,7 @@ class TestMultiPodBatchCostClaim: await asyncio.Event().wait() with self._billing_patches(journal, during_fetch=_never_returns) as logging_obj: - interrupted = asyncio.create_task( - self._instance(prisma, self._router()).check_batch_cost() - ) + interrupted = asyncio.create_task(self._instance(prisma, self._router()).check_batch_cost()) await asyncio.wait_for(reached_fetch.wait(), timeout=5) assert row.batch_processed is False, "an in-flight costing must not mark the row processed" interrupted.cancel() @@ -3363,9 +3145,7 @@ class TestMultiPodBatchCostClaim: await finish_fetch.wait() with self._billing_patches(journal, during_fetch=_wait_for_the_delete_attempt): - costing = asyncio.create_task( - self._instance(prisma, self._router()).check_batch_cost() - ) + costing = asyncio.create_task(self._instance(prisma, self._router()).check_batch_cost()) await asyncio.wait_for(reached_fetch.wait(), timeout=5) with pytest.raises(HTTPException) as blocked: diff --git a/tests/unit/proxy/db/db_transaction_queue/test_daily_spend_update_queue.py b/tests/unit/proxy/db/db_transaction_queue/test_daily_spend_update_queue.py index c17ba75db03..79447d0fc59 100644 --- a/tests/unit/proxy/db/db_transaction_queue/test_daily_spend_update_queue.py +++ b/tests/unit/proxy/db/db_transaction_queue/test_daily_spend_update_queue.py @@ -287,7 +287,7 @@ async def test_queue_max_size_triggers_aggregation( ): """Test that reaching MAX_SIZE_IN_MEMORY_QUEUE triggers aggregation""" # Override MAX_SIZE_IN_MEMORY_QUEUE for testing - litellm._turn_on_debug() + litellm.turn_on_debug() monkeypatch.setattr(daily_spend_update_queue, "MAX_SIZE_IN_MEMORY_QUEUE", 6) test_key = "user1_2023-01-01_key123_gpt-4_openai" diff --git a/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py index 6295469c066..35137e4693e 100644 --- a/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py @@ -241,14 +241,14 @@ class TestHasPostCallGuardrailsForPassthrough: @pytest.mark.asyncio async def test_deferred_flag_stores_and_executes_closure(): """ - When _defer_async_logging is True on logging_obj: + When defer_async_logging is True on logging_obj: 1. wrapper_async stores a callable closure instead of calling create_task 2. Calling the closure fires create_task 3. Sync callbacks fire immediately (not deferred) """ mock_logging_obj = MagicMock() - mock_logging_obj._defer_async_logging = True - mock_logging_obj._enqueue_deferred_logging = None + mock_logging_obj.defer_async_logging = True + mock_logging_obj.enqueue_deferred_logging = None await litellm.acompletion( model="gpt-3.5-turbo", @@ -258,7 +258,7 @@ async def test_deferred_flag_stores_and_executes_closure(): ) # Closure was stored - enqueue_fn = mock_logging_obj._enqueue_deferred_logging + enqueue_fn = mock_logging_obj.enqueue_deferred_logging assert callable(enqueue_fn), "Closure should be stored on logging_obj" # Sync callbacks fired immediately @@ -294,8 +294,8 @@ async def test_deferred_slot_keeps_the_innermost_wrapper_result(): has_logged dedupe keeps the first fired task, so the spend log reads usage from the innermost provider-shaped response and never from an outer wrapper's translation of it.""" logging_obj: Final = MagicMock() - logging_obj._defer_async_logging = True - logging_obj._enqueue_deferred_logging = None + logging_obj.defer_async_logging = True + logging_obj.enqueue_deferred_logging = None logging_obj.async_success_handler = AsyncMock() inner_result: Final = object() outer_result: Final = object() @@ -310,7 +310,7 @@ async def test_deferred_slot_keeps_the_innermost_wrapper_result(): is_litellm_internal_call=False, ) - logging_obj._enqueue_deferred_logging() + logging_obj.enqueue_deferred_logging() await _wait_until(lambda: logging_obj.async_success_handler.await_count > 0) logging_obj.async_success_handler.assert_awaited_once() @@ -370,7 +370,7 @@ async def test_deferred_anthropic_messages_bridged_to_the_responses_api_logs_the function_id="deferred-nested-anthropic-messages", dynamic_async_success_callbacks=[recorder], ) - logging_obj._defer_async_logging = True + logging_obj.defer_async_logging = True response: Final = await litellm.anthropic_messages( model="azure/gpt-5.4-nano", @@ -392,7 +392,7 @@ async def test_deferred_anthropic_messages_bridged_to_the_responses_api_logs_the assert response["usage"]["input_tokens"] == 3 assert response["usage"]["cache_read_input_tokens"] == 7333 - logging_obj._enqueue_deferred_logging() + logging_obj.enqueue_deferred_logging() await _wait_until(lambda: recorder.standard_logging_object is not None) assert recorder.standard_logging_object is not None @@ -408,7 +408,7 @@ async def test_deferred_anthropic_messages_bridged_to_the_responses_api_logs_the @pytest.mark.asyncio async def test_no_flag_fires_create_task_normally(): - """Without _defer_async_logging, wrapper_async calls create_task as before.""" + """Without defer_async_logging, wrapper_async calls create_task as before.""" created_tasks = [] real_create_task = asyncio.create_task @@ -448,7 +448,7 @@ def test_native_pending_logging_is_released_only_for_ocr(call_type: str, excepti logger: Final = MagicMock( call_type=call_type, _native_pending_logging=pending, - _enqueue_deferred_logging=enqueue, + enqueue_deferred_logging=enqueue, ) ProxyBaseLLMRequestProcessing._flush_deferred_async_logging( @@ -484,7 +484,7 @@ def test_flush_deferred_async_logging_fires_on_success(): enqueue_called = True logging_obj = MagicMock() - logging_obj._enqueue_deferred_logging = mock_enqueue + logging_obj.enqueue_deferred_logging = mock_enqueue ProxyBaseLLMRequestProcessing._flush_deferred_async_logging( logging_obj=logging_obj, @@ -492,7 +492,7 @@ def test_flush_deferred_async_logging_fires_on_success(): ) assert enqueue_called is True - assert logging_obj._enqueue_deferred_logging is None + assert logging_obj.enqueue_deferred_logging is None def test_flush_deferred_async_logging_suppressed_on_exception(): @@ -517,7 +517,7 @@ def test_flush_deferred_async_logging_suppressed_on_exception(): enqueue_called = True logging_obj = MagicMock() - logging_obj._enqueue_deferred_logging = mock_enqueue + logging_obj.enqueue_deferred_logging = mock_enqueue ProxyBaseLLMRequestProcessing._flush_deferred_async_logging( logging_obj=logging_obj, @@ -529,7 +529,7 @@ def test_flush_deferred_async_logging_suppressed_on_exception(): "post_call_failure_hook writes its own failure log." ) # Slot is still cleared so a follow-up flush does not double-fire. - assert logging_obj._enqueue_deferred_logging is None + assert logging_obj.enqueue_deferred_logging is None def test_flush_deferred_async_logging_noop_when_no_closure_stored(): @@ -541,7 +541,7 @@ def test_flush_deferred_async_logging_noop_when_no_closure_stored(): class _Bare: pass - logging_obj = _Bare() # no _enqueue_deferred_logging attribute + logging_obj = _Bare() # no enqueue_deferred_logging attribute ProxyBaseLLMRequestProcessing._flush_deferred_async_logging( logging_obj=logging_obj, @@ -553,7 +553,7 @@ def test_flush_deferred_async_logging_noop_when_no_closure_stored(): ) # Helper must not create the attribute as a side effect. - assert not hasattr(logging_obj, "_enqueue_deferred_logging") + assert not hasattr(logging_obj, "enqueue_deferred_logging") def test_proxy_finally_block_routes_through_flush_helper(): @@ -576,11 +576,11 @@ def test_proxy_finally_block_routes_through_flush_helper(): "the request path must call _flush_deferred_async_logging from its " "finally block — do not inline the gating logic." ) - # Belt-and-braces: the inlined `_enqueue_deferred_logging = None` reset + # Belt-and-braces: the inlined `enqueue_deferred_logging = None` reset # was the symptom of the duplicate-log bug; assert it stays inside the # helper, not in the request-processing function. - assert "_enqueue_deferred_logging = None" not in src, ( - "Reset of _enqueue_deferred_logging must live inside " + assert "enqueue_deferred_logging = None" not in src, ( + "Reset of enqueue_deferred_logging must live inside " "_flush_deferred_async_logging, not in the request path." ) @@ -595,14 +595,14 @@ def test_flush_deferred_async_logging_swallows_closure_errors(): raise RuntimeError("logger failure") logging_obj = MagicMock() - logging_obj._enqueue_deferred_logging = boom + logging_obj.enqueue_deferred_logging = boom # Should not raise. ProxyBaseLLMRequestProcessing._flush_deferred_async_logging( logging_obj=logging_obj, exception_raised=False, ) - assert logging_obj._enqueue_deferred_logging is None + assert logging_obj.enqueue_deferred_logging is None # --------------------------------------------------------------------------- @@ -1265,7 +1265,7 @@ class TestDeferredStreamingClosure: ) # litellm_params with no recognized async marker -> classified sync. logging_obj.model_call_details["litellm_params"] = {} - assert LiteLLMLoggingObj._is_sync_litellm_request({}) is True + assert LiteLLMLoggingObj.is_sync_litellm_request({}) is True with ( patch.object( diff --git a/tests/unit/proxy/guardrails/test_guardrail_coverage.py b/tests/unit/proxy/guardrails/test_guardrail_coverage.py index 49c64403313..474ca9d8fa0 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_coverage.py +++ b/tests/unit/proxy/guardrails/test_guardrail_coverage.py @@ -443,10 +443,10 @@ def test_banned_keywords_blocks_multimodal_content(monkeypatch): misaligned with the runtime, so the test wouldn't catch regressions. """ monkeypatch.setattr("litellm.banned_keywords_list", ["forbidden"], raising=False) - from enterprise.enterprise_hooks.banned_keywords import _ENTERPRISE_BannedKeywords + from enterprise.enterprise_hooks.banned_keywords import ENTERPRISE_BannedKeywords from fastapi import HTTPException - guard = _ENTERPRISE_BannedKeywords() + guard = ENTERPRISE_BannedKeywords() async def _run(): await guard.async_pre_call_hook( @@ -475,10 +475,10 @@ def test_banned_keywords_blocks_multimodal_content(monkeypatch): def test_banned_keywords_blocks_responses_api_input(monkeypatch): monkeypatch.setattr("litellm.banned_keywords_list", ["forbidden"], raising=False) - from enterprise.enterprise_hooks.banned_keywords import _ENTERPRISE_BannedKeywords + from enterprise.enterprise_hooks.banned_keywords import ENTERPRISE_BannedKeywords from fastapi import HTTPException - guard = _ENTERPRISE_BannedKeywords() + guard = ENTERPRISE_BannedKeywords() async def _run(): await guard.async_pre_call_hook( @@ -502,10 +502,10 @@ def test_banned_keywords_fires_on_text_content_call_types(monkeypatch, call_type ``acompletion`` (chat completions) and ``aresponses`` (Responses API). """ monkeypatch.setattr("litellm.banned_keywords_list", ["forbidden"], raising=False) - from enterprise.enterprise_hooks.banned_keywords import _ENTERPRISE_BannedKeywords + from enterprise.enterprise_hooks.banned_keywords import ENTERPRISE_BannedKeywords from fastapi import HTTPException - guard = _ENTERPRISE_BannedKeywords() + guard = ENTERPRISE_BannedKeywords() import asyncio @@ -529,9 +529,9 @@ def test_banned_keywords_skips_non_text_call_types(monkeypatch): even when the request body otherwise looks like a chat payload. """ monkeypatch.setattr("litellm.banned_keywords_list", ["forbidden"], raising=False) - from enterprise.enterprise_hooks.banned_keywords import _ENTERPRISE_BannedKeywords + from enterprise.enterprise_hooks.banned_keywords import ENTERPRISE_BannedKeywords - guard = _ENTERPRISE_BannedKeywords() + guard = ENTERPRISE_BannedKeywords() import asyncio @@ -552,10 +552,10 @@ async def test_banned_keywords_post_call_checks_all_choices(monkeypatch, user_ap """Krrish blocker: ``n>1`` responses must not bypass post-call checks by placing the banned text in ``choices[1+]``.""" monkeypatch.setattr("litellm.banned_keywords_list", ["forbidden"], raising=False) - from enterprise.enterprise_hooks.banned_keywords import _ENTERPRISE_BannedKeywords + from enterprise.enterprise_hooks.banned_keywords import ENTERPRISE_BannedKeywords from fastapi import HTTPException - guard = _ENTERPRISE_BannedKeywords() + guard = ENTERPRISE_BannedKeywords() response = ModelResponse( choices=[ Choices(index=0, message=Message(role="assistant", content="clean")), @@ -725,10 +725,10 @@ async def test_openai_moderation_inspects_multimodal_content(monkeypatch, user_a list-format text parts and Responses-API input — without this, multimodal content silently passed moderation.""" from enterprise.enterprise_hooks.openai_moderation import ( - _ENTERPRISE_OpenAI_Moderation, + ENTERPRISE_OpenAI_Moderation, ) - guard = _ENTERPRISE_OpenAI_Moderation() + guard = ENTERPRISE_OpenAI_Moderation() seen_inputs = [] @@ -777,11 +777,11 @@ async def test_openai_moderation_reads_model_name_at_call_time( """``litellm_settings`` applies ``callbacks`` and ``openai_moderations_model_name`` in YAML order, so the hook must resolve the model when it runs, not when it is constructed.""" from enterprise.enterprise_hooks.openai_moderation import ( - _ENTERPRISE_OpenAI_Moderation, + ENTERPRISE_OpenAI_Moderation, ) monkeypatch.setattr(litellm, "openai_moderations_model_name", None) - guard = _ENTERPRISE_OpenAI_Moderation() + guard = ENTERPRISE_OpenAI_Moderation() monkeypatch.setattr(litellm, "openai_moderations_model_name", configured_after_init) class FakeModeration: @@ -808,10 +808,10 @@ async def test_google_text_moderation_inspects_multimodal_content(user_api_key): """The text passed to Google's moderation client must include list-format text parts.""" from enterprise.enterprise_hooks.google_text_moderation import ( - _ENTERPRISE_GoogleTextModeration, + ENTERPRISE_GoogleTextModeration, ) - guard = _ENTERPRISE_GoogleTextModeration.__new__(_ENTERPRISE_GoogleTextModeration) + guard = ENTERPRISE_GoogleTextModeration.__new__(ENTERPRISE_GoogleTextModeration) seen_documents = [] def fake_language_document(content, type_): diff --git a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py index 2b117f05b92..7397b0ec6e2 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py @@ -1646,7 +1646,7 @@ async def test_get_guardrail_info_endpoint_config_guardrail(mocker): # Mock _get_masked_values to return values as-is mocker.patch( - "litellm.litellm_core_utils.litellm_logging._get_masked_values", + "litellm.litellm_core_utils.litellm_logging.get_masked_values", side_effect=lambda x, **kwargs: x, ) diff --git a/tests/unit/proxy/guardrails/test_guardrail_registry.py b/tests/unit/proxy/guardrails/test_guardrail_registry.py index 2f29f964e63..2f171677d12 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_registry.py +++ b/tests/unit/proxy/guardrails/test_guardrail_registry.py @@ -641,7 +641,7 @@ def test_repeated_db_sync_does_not_accumulate_runner_instances(): def distinct_runner_instances() -> int: seen = set() - for callback in litellm.logging_callback_manager._get_all_callbacks(): + for callback in litellm.logging_callback_manager.get_all_callbacks(): if isinstance(callback, CustomGuardrail) and getattr(callback, "guardrail_name", None) == name: seen.add(id(callback)) return len(seen) diff --git a/tests/unit/proxy/hooks/test_banned_keyword_list.py b/tests/unit/proxy/hooks/test_banned_keyword_list.py index 82ab693e699..81c97b66644 100644 --- a/tests/unit/proxy/hooks/test_banned_keyword_list.py +++ b/tests/unit/proxy/hooks/test_banned_keyword_list.py @@ -12,7 +12,7 @@ load_dotenv() import pytest import litellm from litellm.proxy.enterprise.enterprise_hooks.banned_keywords import ( - _ENTERPRISE_BannedKeywords, + ENTERPRISE_BannedKeywords, ) from litellm import Router, mock_completion from litellm.proxy.utils import ProxyLogging, hash_token @@ -29,7 +29,7 @@ async def test_banned_keywords_check(): """ litellm.banned_keywords_list = ["hello"] - banned_keywords_obj = _ENTERPRISE_BannedKeywords() + banned_keywords_obj = ENTERPRISE_BannedKeywords() _api_key = "sk-98765" _api_key = hash_token("sk-98765") diff --git a/tests/unit/proxy/hooks/test_batch_file_validation.py b/tests/unit/proxy/hooks/test_batch_file_validation.py index 4ef94f0965b..38ee6997899 100644 --- a/tests/unit/proxy/hooks/test_batch_file_validation.py +++ b/tests/unit/proxy/hooks/test_batch_file_validation.py @@ -31,9 +31,9 @@ def _models(file_content_as_dict): def test_token_counter_counts_chat_messages(): - from litellm.batches.batch_utils import _count_entry_tokens + from litellm.batches.batch_utils import count_entry_tokens - tokens = _count_entry_tokens( + tokens = count_entry_tokens( { "body": { "model": "gpt-4o-mini", @@ -47,27 +47,27 @@ def test_token_counter_counts_chat_messages(): def test_token_counter_counts_text_completion_prompt(): """Pre-fix this returned 0 tokens (the counter only inspected `messages`), letting `prompt`-style batches slip past TPM limits.""" - from litellm.batches.batch_utils import _count_entry_tokens + from litellm.batches.batch_utils import count_entry_tokens - tokens = _count_entry_tokens( + tokens = count_entry_tokens( {"body": {"model": "gpt-3.5-turbo-instruct", "prompt": "hello world"}} ) assert tokens > 0 def test_token_counter_counts_embedding_input_string(): - from litellm.batches.batch_utils import _count_entry_tokens + from litellm.batches.batch_utils import count_entry_tokens - tokens = _count_entry_tokens( + tokens = count_entry_tokens( {"body": {"model": "text-embedding-3-small", "input": "hello world"}} ) assert tokens > 0 def test_token_counter_counts_embedding_input_list(): - from litellm.batches.batch_utils import _count_entry_tokens + from litellm.batches.batch_utils import count_entry_tokens - tokens = _count_entry_tokens( + tokens = count_entry_tokens( { "body": { "model": "text-embedding-3-small", @@ -79,9 +79,9 @@ def test_token_counter_counts_embedding_input_list(): def test_token_counter_counts_text_completion_prompt_list(): - from litellm.batches.batch_utils import _count_entry_tokens + from litellm.batches.batch_utils import count_entry_tokens - tokens = _count_entry_tokens( + tokens = count_entry_tokens( { "body": { "model": "gpt-3.5-turbo-instruct", @@ -96,9 +96,9 @@ def test_token_counter_counts_pre_tokenized_prompt_int_list(): """OpenAI's text-completion API accepts a single pre-tokenized prompt as a list of ints. Each int is one token; pre-fix this shape was silently counted as zero, leaving a TPM bypass.""" - from litellm.batches.batch_utils import _count_entry_tokens + from litellm.batches.batch_utils import count_entry_tokens - tokens = _count_entry_tokens( + tokens = count_entry_tokens( { "body": { "model": "gpt-3.5-turbo-instruct", @@ -113,9 +113,9 @@ def test_token_counter_counts_pre_tokenized_prompt_list_of_int_lists(): """Multiple pre-tokenized prompts (`list[list[int]]`) — the most important bypass shape. A 1000-token batch must report 1000 tokens, not zero.""" - from litellm.batches.batch_utils import _count_entry_tokens + from litellm.batches.batch_utils import count_entry_tokens - tokens = _count_entry_tokens( + tokens = count_entry_tokens( { "body": { "model": "gpt-3.5-turbo-instruct", @@ -128,9 +128,9 @@ def test_token_counter_counts_pre_tokenized_prompt_list_of_int_lists(): def test_token_counter_counts_pre_tokenized_input_for_embeddings(): """Same shape applies to embeddings (`input`).""" - from litellm.batches.batch_utils import _count_entry_tokens + from litellm.batches.batch_utils import count_entry_tokens - tokens = _count_entry_tokens( + tokens = count_entry_tokens( { "body": { "model": "text-embedding-3-small", @@ -1741,13 +1741,13 @@ def _make_batch_input_bytes(n_rows: int, padding: int = 200) -> bytes: def test_iter_batch_output_entries_matches_dict_list(): from litellm.batches.batch_utils import ( - _get_file_content_as_dictionary, + get_file_content_as_dictionary, _iter_batch_output_entries, ) raw = _make_batch_input_bytes(50) streamed = list(_iter_batch_output_entries(raw)) - assert streamed == _get_file_content_as_dictionary(raw) + assert streamed == get_file_content_as_dictionary(raw) assert streamed[0]["custom_id"] == "request-0" # tolerant of blank lines and a missing trailing newline assert list(_iter_batch_output_entries(raw + b"\n\n")) == streamed @@ -1758,7 +1758,7 @@ def test_streaming_count_peak_below_dict_list(): import tracemalloc from litellm.batches.batch_utils import ( - _get_file_content_as_dictionary, + get_file_content_as_dictionary, _iter_batch_output_entries, ) @@ -1785,7 +1785,7 @@ def test_streaming_count_peak_below_dict_list(): return count def _build_list(): - return len(_get_file_content_as_dictionary(raw)) + return len(get_file_content_as_dictionary(raw)) stream_peak = _measure(_stream) list_peak = _measure(_build_list) @@ -1813,7 +1813,7 @@ async def test_count_input_file_usage_streams_without_building_list(): with ( patch("litellm.afile_content", new=AsyncMock(return_value=fake_content)), patch( - "litellm.batches.batch_utils._get_file_content_as_dictionary" + "litellm.batches.batch_utils.get_file_content_as_dictionary" ) as mock_dict_list, ): usage = await rate_limiter.count_input_file_usage( @@ -1874,7 +1874,7 @@ async def test_count_input_file_usage_enforces_models_when_token_counting_fails( with ( patch("litellm.afile_content", new=AsyncMock(return_value=fake_content)), - patch("litellm.proxy.hooks.batch_rate_limiter._count_entry_tokens", new=_boom), + patch("litellm.proxy.hooks.batch_rate_limiter.count_entry_tokens", new=_boom), patch("litellm.proxy.auth.auth_checks.can_key_call_model", new=deny), patch("litellm.proxy.proxy_server.llm_router", MagicMock(model_list=[])), ): @@ -1919,7 +1919,7 @@ async def test_count_input_file_usage_estimates_tokens_when_counting_fails_for_a with ( patch("litellm.afile_content", new=AsyncMock(return_value=fake_content)), - patch("litellm.proxy.hooks.batch_rate_limiter._count_entry_tokens", new=_boom), + patch("litellm.proxy.hooks.batch_rate_limiter.count_entry_tokens", new=_boom), patch("litellm.proxy.auth.auth_checks.can_key_call_model", new=allow), patch("litellm.proxy.proxy_server.llm_router", MagicMock(model_list=[])), ): diff --git a/tests/unit/proxy/hooks/test_proxy_hooks_init.py b/tests/unit/proxy/hooks/test_proxy_hooks_init.py index a6edd3db944..7f07fa7966c 100644 --- a/tests/unit/proxy/hooks/test_proxy_hooks_init.py +++ b/tests/unit/proxy/hooks/test_proxy_hooks_init.py @@ -23,14 +23,14 @@ def test_managed_files_hook_registered(): pytest.importorskip("litellm_enterprise") assert "managed_files" in PROXY_HOOKS hook_cls = get_proxy_hook("managed_files") - assert hook_cls.__name__ == "_PROXY_LiteLLMManagedFiles" + assert hook_cls.__name__ == "PROXY_LiteLLMManagedFiles" def test_managed_vector_stores_hook_registered(): pytest.importorskip("litellm_enterprise") assert "managed_vector_stores" in PROXY_HOOKS hook_cls = get_proxy_hook("managed_vector_stores") - assert hook_cls.__name__ == "_PROXY_LiteLLMManagedVectorStores" + assert hook_cls.__name__ == "PROXY_LiteLLMManagedVectorStores" def test_isolation_module_does_not_pull_in_proxy_utils(): diff --git a/tests/unit/proxy/image_endpoints/test_azure_routes.py b/tests/unit/proxy/image_endpoints/test_azure_routes.py index 46fe9a6f893..87cb08e9176 100644 --- a/tests/unit/proxy/image_endpoints/test_azure_routes.py +++ b/tests/unit/proxy/image_endpoints/test_azure_routes.py @@ -100,7 +100,7 @@ def test_azure_image_generation_route(client_no_auth): def test_azure_image_edit_route(client_no_auth): - litellm._turn_on_debug() + litellm.turn_on_debug() client, _, mock_aimage_edit = client_no_auth image_path = os.path.join( os.path.dirname(__file__), diff --git a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py index 4862af63f12..4439a3511df 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -1761,7 +1761,7 @@ class TestTeamModelSiblingRouting: ) # Global deployment should be accessible when team_id is provided - deployments = router._get_all_deployments(model_name="global-gpt-4o", team_id="teamA") + deployments = router.get_all_deployments(model_name="global-gpt-4o", team_id="teamA") assert len(deployments) == 1 assert deployments[0]["model_name"] == "global-gpt-4o" diff --git a/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py index 25c62fc6dcc..4ce31cbb449 100644 --- a/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py @@ -4849,7 +4849,7 @@ def _setup_unscoped_list_files_route_over_real_hook( import litellm.proxy.proxy_server as ps from litellm.proxy._types import LitellmUserRoles from litellm_enterprise.proxy.hooks.managed_files import ( - _PROXY_LiteLLMManagedFiles, + PROXY_LiteLLMManagedFiles, ) for env_var in ("OPENAI_API_KEY", "OPENAI_ADMIN_KEY", "OPENAI_ORGANIZATION"): @@ -4857,7 +4857,7 @@ def _setup_unscoped_list_files_route_over_real_hook( monkeypatch.setattr(litellm, "api_key", None, raising=False) monkeypatch.setattr(litellm, "openai_key", None, raising=False) - managed_files = _PROXY_LiteLLMManagedFiles( + managed_files = PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), prisma_client=MagicMock() ) managed_files.prisma_client.db.litellm_managedfiletable = _ManagedFileTableOverRows(rows) diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py index 3d42301ed11..eeac2c9f02d 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py @@ -1684,8 +1684,8 @@ class TestOpenAIPassthroughIntegration: ) image_response._hidden_params = {"response_cost": test_cost} - # Test the _response_cost_calculator method - calculated_cost = logging_obj._response_cost_calculator(result=image_response) + # Test the response_cost_calculator method + calculated_cost = logging_obj.response_cost_calculator(result=image_response) assert calculated_cost == test_cost, f"Expected {test_cost}, got {calculated_cost}" diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 24ea1fe8d53..79b0828ec15 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -1140,7 +1140,7 @@ class TestVertexAIPassThroughHandler: end_time=end_time, cache_hit=False, ) - recomputed: Final = logging_obj._response_cost_calculator(result=result["result"]) + recomputed: Final = logging_obj.response_cost_calculator(result=result["result"]) return result["kwargs"]["response_cost"], recomputed global_handler_cost, global_recomputed_cost = costs_for("global") diff --git a/tests/unit/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py b/tests/unit/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py index 00a3606bbf5..e6b19f4eec1 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py +++ b/tests/unit/proxy/pass_through_endpoints/test_streaming_handler_interrupt.py @@ -261,7 +261,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker_on_disconne def _logging_obj_with_write_once_cst(): - """Build a MagicMock that mirrors the real Logging behavior: _update_completion_start_time + """Build a MagicMock that mirrors the real Logging behavior: update_completion_start_time latches self.completion_start_time so the write-once guard actually latches.""" obj = _unarmed_logging_obj() obj.completion_start_time = None @@ -269,7 +269,7 @@ def _logging_obj_with_write_once_cst(): def _update(*, completion_start_time): obj.completion_start_time = completion_start_time - obj._update_completion_start_time.side_effect = _update + obj.update_completion_start_time.side_effect = _update return obj @@ -305,8 +305,8 @@ async def test_chunk_processor_stamps_completion_start_time_on_first_chunk(): await asyncio.sleep(0) assert received == chunks - mock_logging_obj._update_completion_start_time.assert_called_once() - stamped = mock_logging_obj._update_completion_start_time.call_args.kwargs["completion_start_time"] + mock_logging_obj.update_completion_start_time.assert_called_once() + stamped = mock_logging_obj.update_completion_start_time.call_args.kwargs["completion_start_time"] assert isinstance(stamped, datetime) @@ -340,7 +340,7 @@ async def test_chunk_processor_does_not_reset_completion_start_time_on_later_chu ): pass - mock_logging_obj._update_completion_start_time.assert_not_called() + mock_logging_obj.update_completion_start_time.assert_not_called() assert mock_logging_obj.completion_start_time == real_first @@ -377,7 +377,7 @@ async def test_chunk_processor_stamps_completion_start_time_on_cost_injection_pa finally: litellm_mod.include_cost_in_streaming_usage = original - mock_logging_obj._update_completion_start_time.assert_called_once() + mock_logging_obj.update_completion_start_time.assert_called_once() def _openai_passthrough_stream_chunks(): diff --git a/tests/unit/proxy/proxy_server/test_background_health.py b/tests/unit/proxy/proxy_server/test_background_health.py index d99aa637955..b15349d6705 100644 --- a/tests/unit/proxy/proxy_server/test_background_health.py +++ b/tests/unit/proxy/proxy_server/test_background_health.py @@ -433,7 +433,7 @@ def test_write_health_state_to_router_cache_sets_states(monkeypatch): import litellm.router_utils.cooldown_handlers as cd - monkeypatch.setattr(cd, "_set_cooldown_deployments", lambda **_kw: None) + monkeypatch.setattr(cd, "set_cooldown_deployments", lambda **_kw: None) import litellm.router_utils.router_callbacks.track_deployment_metrics as tdm @@ -507,7 +507,7 @@ def test_write_health_state_to_router_cache_populates_for_listing_filter(monkeyp monkeypatch.setattr( cd, - "_set_cooldown_deployments", + "set_cooldown_deployments", lambda **kw: cooldowns.append(kw.get("deployment")), ) @@ -560,7 +560,7 @@ def test_write_health_state_to_router_cache_swallows_internal_failures(monkeypat @pytest.mark.asyncio async def test_adaptive_router_flusher_loop_flushes_each_router(monkeypatch): fake_ar = MagicMock() - fake_ar._state_loaded = True + fake_ar.state_loaded = True fake_ar.queue.flush_state_to_db = AsyncMock() fake_ar.queue.flush_session_to_db = AsyncMock() diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index 0b58e5d141f..47f17a10969 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -2976,7 +2976,7 @@ def test_ProxyConfig__add_deployment_pinned_row_follows_the_cost_map_across_relo assert ProxyConfig()._add_deployment(db_models=[pinned, typed]) == 2 monkeypatch.setitem(litellm.model_cost["gpt-5.6"], "input_cost_per_token", 1e-06) - router._replay_model_cost_registrations() + router.replay_model_cost_registrations() assert litellm.model_cost.get("pinned-row", {}).get("input_cost_per_token") is None assert router.get_deployment(model_id="pinned-row").model_info.input_cost_per_token is None @@ -3008,7 +3008,7 @@ def test_ProxyConfig__add_deployment_ptu_row_with_a_cost_map_copy_still_bills_ze ) assert ProxyConfig()._add_deployment(db_models=[ptu]) == 1 - router._replay_model_cost_registrations() + router.replay_model_cost_registrations() assert litellm.model_cost["ptu-row"]["input_cost_per_token"] == 0.0 assert litellm.model_cost["ptu-row"]["output_cost_per_token"] == 0.0 diff --git a/tests/unit/proxy/proxy_server/test_routes_config.py b/tests/unit/proxy/proxy_server/test_routes_config.py index 4234cdad23d..c470fa96d42 100644 --- a/tests/unit/proxy/proxy_server/test_routes_config.py +++ b/tests/unit/proxy/proxy_server/test_routes_config.py @@ -1508,7 +1508,7 @@ def test_get_config_callbacks_excludes_internal_runtime_callbacks(client, auth_a ) monkeypatch.setattr(ps, "proxy_config", fake_proxy_config) - from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles import litellm from litellm._service_logger import ServiceLogging @@ -1547,7 +1547,7 @@ def test_get_config_callbacks_excludes_internal_runtime_callbacks(client, auth_a "callbacks", [ _PROXY_CacheControlCheck(), - _PROXY_LiteLLMManagedFiles(internal_usage_cache=MagicMock(), prisma_client=MagicMock()), + PROXY_LiteLLMManagedFiles(internal_usage_cache=MagicMock(), prisma_client=MagicMock()), ServiceLogging(), VectorStorePreCallHook(), _InventoryTestGuardrail(guardrail_name="inventory-test-guardrail"), diff --git a/tests/unit/proxy/proxy_server/test_routes_utils.py b/tests/unit/proxy/proxy_server/test_routes_utils.py index 91329d122ee..295b9e5ae8b 100644 --- a/tests/unit/proxy/proxy_server/test_routes_utils.py +++ b/tests/unit/proxy/proxy_server/test_routes_utils.py @@ -33,7 +33,7 @@ def patched_token_counter(monkeypatch): monkeypatch.setattr(litellm, "disable_token_counter", False, raising=False) monkeypatch.setattr( litellm.utils, - "_select_tokenizer", + "select_tokenizer", lambda model, custom_tokenizer=None: { "type": "openai_tokenizer", "tokenizer": None, @@ -255,10 +255,10 @@ def lookup_fixture_model(monkeypatch): monkeypatch.setattr(proxy_server, "llm_router", None) monkeypatch.setitem(litellm.model_cost, "lookup-fixture-model", entry) litellm.get_model_info.cache_clear() - litellm.utils._cached_get_model_info_helper.cache_clear() + litellm.utils.cached_get_model_info_helper.cache_clear() yield entry litellm.get_model_info.cache_clear() - litellm.utils._cached_get_model_info_helper.cache_clear() + litellm.utils.cached_get_model_info_helper.cache_clear() def test_model_info_lookup_returns_full_cost_map_entry_for_unregistered_model(client, auth_as, lookup_fixture_model): diff --git a/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py index 83c3441dc17..2e3b3c13bb3 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py @@ -3685,7 +3685,7 @@ class TestSpendLogsPayload: @pytest.mark.asyncio async def test_spend_logs_payload_e2e(self): litellm.callbacks = [_ProxyDBLogger(message_logging=False)] - # litellm._turn_on_debug() + # litellm.turn_on_debug() with ( patch.object( diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 8dc6c0477b0..368f55b1eb1 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -3721,7 +3721,7 @@ class TestStreamingOverheadHeader: mock_logging_obj.caching_details = None mock_logging_obj.callback_duration_ms = None mock_logging_obj.litellm_call_id = "test-call-id" - mock_logging_obj._response_cost_calculator = MagicMock(return_value=0.001) + mock_logging_obj.response_cost_calculator = MagicMock(return_value=0.001) # Simulate a streaming result object with _hidden_params (like CustomStreamWrapper) stream_result = MagicMock() @@ -4857,7 +4857,7 @@ class TestDisconnectGatherCleanup: mock_logging_obj = MagicMock() mock_logging_obj.litellm_call_id = "test-call-id" - mock_logging_obj._defer_async_logging = False + mock_logging_obj.defer_async_logging = False mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.during_call_hook = AsyncMock(return_value=None) @@ -4904,7 +4904,7 @@ class TestDisconnectGatherCleanup: mock_logging_obj = MagicMock() mock_logging_obj.litellm_call_id = "test-call-id" - mock_logging_obj._defer_async_logging = False + mock_logging_obj.defer_async_logging = False mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.during_call_hook = AsyncMock(return_value=None) @@ -4967,7 +4967,7 @@ class TestDisconnectGatherCleanup: mock_logging_obj = MagicMock() mock_logging_obj.litellm_call_id = "test-call-id" - mock_logging_obj._defer_async_logging = False + mock_logging_obj.defer_async_logging = False mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.during_call_hook = slow_during_call_hook @@ -5051,7 +5051,7 @@ class TestDisconnectGatherCleanup: mock_logging_obj = MagicMock() mock_logging_obj.litellm_call_id = "test-call-id" - mock_logging_obj._defer_async_logging = False + mock_logging_obj.defer_async_logging = False mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.during_call_hook = successful_hook @@ -5103,7 +5103,7 @@ async def test_response_model_echoes_the_name_the_client_sent_before_auth_rewrot async def fake_route_request(**_kwargs): return llm() - logging_obj = MagicMock(litellm_call_id="call-id", _defer_async_logging=False) + logging_obj = MagicMock(litellm_call_id="call-id", defer_async_logging=False) proxy_logging = MagicMock(spec=ProxyLogging) proxy_logging.during_call_hook = AsyncMock(return_value=None) proxy_logging.post_call_success_hook = AsyncMock(side_effect=lambda data, user_api_key_dict, response: response) @@ -5456,7 +5456,7 @@ class TestCancelOnDisconnect: logging_obj = MagicMock() logging_obj.litellm_call_id = "test-cancel-on-disconnect" - logging_obj._defer_async_logging = False + logging_obj.defer_async_logging = False logging_obj._on_deferred_stream_complete = None logging_obj.cost_breakdown = None @@ -6125,8 +6125,8 @@ class TestResponseCostHeaderForTypedDictResponses: logging_obj.litellm_call_id = "call-lit4076" logging_obj.cost_breakdown = None logging_obj.model_call_details = model_call_details - logging_obj._response_cost_calculator = response_cost_calculator - logging_obj._enqueue_deferred_logging = None + logging_obj.response_cost_calculator = response_cost_calculator + logging_obj.enqueue_deferred_logging = None logging_obj._on_deferred_stream_complete = None return logging_obj @@ -6314,7 +6314,7 @@ class TestResponseCostHeaderForTypedDictResponses: logging_obj = self._build_logging_obj( model_call_details={}, - response_cost_calculator=real_logging._response_cost_calculator, + response_cost_calculator=real_logging.response_cost_calculator, ) fastapi_response = await self._drive_non_streaming( @@ -6518,8 +6518,8 @@ class TestCostHeadersForCallsPricedAtZero: logging_obj.litellm_params = {} logging_obj.cost_breakdown = None logging_obj.model_call_details = {"response_cost": recovered_cost} - logging_obj._response_cost_calculator = MagicMock(return_value=recovered_cost) - logging_obj._enqueue_deferred_logging = None + logging_obj.response_cost_calculator = MagicMock(return_value=recovered_cost) + logging_obj.enqueue_deferred_logging = None logging_obj._on_deferred_stream_complete = None return logging_obj @@ -8466,7 +8466,7 @@ class TestInjectCostIntoUsageDict: self._cost = cost self.captured_result = None - def _response_cost_calculator(self, result): + def response_cost_calculator(self, result): self.captured_result = result return self._cost @@ -8498,7 +8498,7 @@ class TestInjectCostIntoUsageDict: def test_message_delta_falls_back_to_model_pricing_when_the_logging_obj_returns_no_cost(self): class _StubLoggingObj: - def _response_cost_calculator(self, result): + def response_cost_calculator(self, result): return None model = "claude-haiku-4-5" @@ -8523,7 +8523,7 @@ class TestInjectCostIntoUsageDict: model-name pricing rather than propagating into the response body.""" class _StubLoggingObj: - def _response_cost_calculator(self, result): + def response_cost_calculator(self, result): raise ValueError("no pricing for this deployment") model = "claude-haiku-4-5" @@ -8612,7 +8612,7 @@ class TestInjectCostIntoUsageDict: self._cost = cost self.captured_result = None - def _response_cost_calculator(self, result): + def response_cost_calculator(self, result): self.captured_result = result return self._cost @@ -8636,7 +8636,7 @@ class TestInjectCostIntoUsageDict: def test_openai_chunk_falls_back_to_model_pricing_when_the_logging_obj_returns_no_cost(self): class _StubLoggingObj: - def _response_cost_calculator(self, result): + def response_cost_calculator(self, result): return None event = { @@ -8691,7 +8691,7 @@ class TestProcessChunkWithCostInjection: monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True) class _StubLoggingObj: - def _response_cost_calculator(self, result): + def response_cost_calculator(self, result): return 0.00042 chunk = ( @@ -9303,9 +9303,9 @@ class TestDetachedStreamFailureHook: logging_obj = MagicMock() logging_obj.litellm_call_id = "call-lit3798" logging_obj.model_call_details = {} - logging_obj._enqueue_deferred_logging = None + logging_obj.enqueue_deferred_logging = None logging_obj._on_deferred_stream_complete = None - logging_obj._on_detached_stream_failure = None + logging_obj.on_detached_stream_failure = None return logging_obj @staticmethod @@ -9354,7 +9354,7 @@ class TestDetachedStreamFailureHook: ) failure = RuntimeError("upstream died after the client left") - await logging_obj._on_detached_stream_failure(failure) + await logging_obj.on_detached_stream_failure(failure) assert recorder.calls == [ { @@ -9378,7 +9378,7 @@ class TestDetachedStreamFailureHook: ) failure = RuntimeError("upstream died after the client left") - await logging_obj._on_detached_stream_failure(failure) + await logging_obj.on_detached_stream_failure(failure) assert [call["original_exception"] for call in recorder.calls] == [failure] @@ -9391,12 +9391,12 @@ class TestPostCallMaskedOutputReachesDeferredLogging: logging_obj = MagicMock() logging_obj.litellm_call_id = "lit-8325-call" - logging_obj._defer_async_logging = False + logging_obj.defer_async_logging = False logging_obj._on_deferred_stream_complete = None logging_obj.cost_breakdown = None logging_obj.model_call_details = {} recorded_at_enqueue: dict[str, object] = {} - logging_obj._enqueue_deferred_logging = lambda: recorded_at_enqueue.update(logging_obj.model_call_details) + logging_obj.enqueue_deferred_logging = lambda: recorded_at_enqueue.update(logging_obj.model_call_details) processor = ProxyBaseLLMRequestProcessing(data={"model": "oa", "litellm_logging_obj": logging_obj}) @@ -9478,7 +9478,7 @@ class TestStreamingResponseHeadersFollowFallback: logging_obj = MagicMock() logging_obj.litellm_call_id = "lit-6767-call" - logging_obj._defer_async_logging = False + logging_obj.defer_async_logging = False logging_obj._on_deferred_stream_complete = None logging_obj.cost_breakdown = None @@ -9541,7 +9541,7 @@ class TestStreamingResponseHeadersFollowFallback: logging_obj = MagicMock() logging_obj.litellm_call_id = "lit-7144-call" - logging_obj._defer_async_logging = False + logging_obj.defer_async_logging = False logging_obj._on_deferred_stream_complete = None logging_obj.cost_breakdown = None processor_data["litellm_logging_obj"] = logging_obj @@ -9595,7 +9595,7 @@ class TestStreamingResponseHeadersFollowFallback: logging_obj = MagicMock() logging_obj.litellm_call_id = "lit-8302-call" - logging_obj._defer_async_logging = False + logging_obj.defer_async_logging = False logging_obj._on_deferred_stream_complete = None logging_obj.cost_breakdown = None processor = ProxyBaseLLMRequestProcessing( @@ -9677,7 +9677,7 @@ async def test_messages_http_headers_refresh_after_lazy_fallback(monkeypatch: py stream = _MessagesFallbackStream() logging_obj = MagicMock() logging_obj.litellm_call_id = "messages-fallback-headers" - logging_obj._defer_async_logging = False + logging_obj.defer_async_logging = False logging_obj._on_deferred_stream_complete = None logging_obj.cost_breakdown = None logging_obj.litellm_params = {} diff --git a/tests/unit/proxy/test_litellm_pre_call_utils.py b/tests/unit/proxy/test_litellm_pre_call_utils.py index e604604025a..d4624719565 100644 --- a/tests/unit/proxy/test_litellm_pre_call_utils.py +++ b/tests/unit/proxy/test_litellm_pre_call_utils.py @@ -2984,7 +2984,7 @@ async def test_add_litellm_metadata_from_request_headers(): Relevant issue: https://github.com/BerriAI/litellm/issues/14008 """ # Set up test logger - litellm._turn_on_debug() + litellm.turn_on_debug() test_logger = TestCustomLogger() original_callbacks = litellm.callbacks litellm.callbacks = [test_logger] @@ -3095,7 +3095,7 @@ async def test_anthropic_messages_standard_logging_object_matches_fixture(): Regression: /v1/messages calls routed to non-Anthropic providers should keep call_type=anthropic_messages in standard logging payloads. """ - litellm._turn_on_debug() + litellm.turn_on_debug() test_logger = TestCustomLogger() original_callbacks = litellm.callbacks litellm.callbacks = [test_logger] diff --git a/tests/unit/proxy/test_prompt_test_endpoint.py b/tests/unit/proxy/test_prompt_test_endpoint.py index 723f6c19c97..ab4614cc314 100644 --- a/tests/unit/proxy/test_prompt_test_endpoint.py +++ b/tests/unit/proxy/test_prompt_test_endpoint.py @@ -27,7 +27,7 @@ User: Hello {{name}}, how are you?""" # Parse the dotprompt prompt_manager = PromptManager() - frontmatter, template_content = prompt_manager._parse_frontmatter( + frontmatter, template_content = prompt_manager.parse_frontmatter( content=dotprompt_content ) @@ -126,7 +126,7 @@ temperature: 0.7 User: Hello""" prompt_manager = PromptManager() - frontmatter, _ = prompt_manager._parse_frontmatter(content=dotprompt_content) + frontmatter, _ = prompt_manager.parse_frontmatter(content=dotprompt_content) model = frontmatter.get("model") assert model is None diff --git a/tests/unit/proxy/test_proxy_config_unit_test.py b/tests/unit/proxy/test_proxy_config_unit_test.py index 2181c932586..bee3c3a9eaa 100644 --- a/tests/unit/proxy/test_proxy_config_unit_test.py +++ b/tests/unit/proxy/test_proxy_config_unit_test.py @@ -240,7 +240,7 @@ def test_add_callbacks_invalid_input(): @pytest.mark.asyncio async def test_json_logs_calls_turn_on_json(): """ - Test that json_logs: true in litellm_settings calls litellm._turn_on_json() + Test that json_logs: true in litellm_settings calls litellm.turn_on_json() This is a regression test for the bug where json_logs in config file would only set the attribute but not actually enable JSON logging. @@ -269,7 +269,7 @@ async def test_json_logs_calls_turn_on_json(): proxy_config = ProxyConfig() # Mock _turn_on_json to track if it gets called - with mock.patch("litellm._turn_on_json") as mock_turn_on_json: + with mock.patch("litellm.turn_on_json") as mock_turn_on_json: await proxy_config.load_config( router=None, config_file_path=temp_file_path, diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index afdca207e24..ec91f1d8edc 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -12071,7 +12071,7 @@ def _reset_runtime_callbacks(monkeypatch: pytest.MonkeyPatch) -> None: def _runtime_callback_names() -> frozenset[str]: manager = litellm.logging_callback_manager - return frozenset(manager._get_callback_string(callback) for callback in manager._get_all_callbacks()) + return frozenset(manager._get_callback_string(callback) for callback in manager.get_all_callbacks()) @pytest.mark.parametrize("setting_key", ["success_callback", "failure_callback", "callbacks"]) @@ -12094,10 +12094,10 @@ def test_db_config_sync_unregisters_a_callback_the_stored_config_no_longer_lists def test_db_config_sync_keeps_callbacks_it_did_not_register(monkeypatch: pytest.MonkeyPatch): import litellm.proxy.proxy_server as ps - from litellm.utils import _add_custom_logger_callback_to_specific_event + from litellm.utils import add_custom_logger_callback_to_specific_event _reset_runtime_callbacks(monkeypatch) - _add_custom_logger_callback_to_specific_event("langfuse_otel", "success") + add_custom_logger_callback_to_specific_event("langfuse_otel", "success") litellm.logging_callback_manager.add_litellm_success_callback("helicone") pc = ps.ProxyConfig() @@ -12140,7 +12140,7 @@ def test_db_config_sync_registers_otel_v2_arize_next_to_otel( import litellm.proxy.proxy_server as ps from litellm.integrations.otel.logger import OpenTelemetryV2 from litellm.integrations.otel.model.config import is_otel_v2_enabled - from litellm.utils import _add_custom_logger_callback_to_specific_event + from litellm.utils import add_custom_logger_callback_to_specific_event _reset_runtime_callbacks(monkeypatch) for extra_list in ("input_callback", "service_callback"): @@ -12154,7 +12154,7 @@ def test_db_config_sync_registers_otel_v2_arize_next_to_otel( is_otel_v2_enabled.cache_clear() try: getattr(litellm.logging_callback_manager, f"add_litellm_{event}_callback")("helicone") - _add_custom_logger_callback_to_specific_event("otel", event) + add_custom_logger_callback_to_specific_event("otel", event) pc = ps.ProxyConfig() for _ in range(2): pc._add_callbacks_from_db_config({"litellm_settings": {setting_key: ["arize"]}}) @@ -15342,7 +15342,7 @@ async def test_token_counter_loads_a_custom_tokenizer_off_the_event_loop(monkeyp from litellm import Router from tests.unit.litellm_core_utils.event_loop_lag import assert_loop_stayed_free, timed_with_loop_lags - claude_tokenizer: Final = litellm.utils._select_tokenizer("claude-fable-5")["tokenizer"] + claude_tokenizer: Final = litellm.utils.select_tokenizer("claude-fable-5")["tokenizer"] class SlowHubTokenizer: @staticmethod @@ -15383,7 +15383,7 @@ async def test_token_counter_loads_a_custom_tokenizer_once_per_identifier_revisi from litellm import Router from litellm.types.router import DeploymentTypedDict - claude_tokenizer: Final[Tokenizer] = litellm.utils._select_tokenizer("claude-fable-5")["tokenizer"] + claude_tokenizer: Final[Tokenizer] = litellm.utils.select_tokenizer("claude-fable-5")["tokenizer"] from_pretrained: Final = MagicMock(return_value=claude_tokenizer) def deployment(model_name: str, revision: str, auth_token: str | None) -> DeploymentTypedDict: diff --git a/tests/unit/proxy/test_proxy_setting_guardrails.py b/tests/unit/proxy/test_proxy_setting_guardrails.py index 71b7783f5ee..c1d2c640b93 100644 --- a/tests/unit/proxy/test_proxy_setting_guardrails.py +++ b/tests/unit/proxy/test_proxy_setting_guardrails.py @@ -62,5 +62,5 @@ def test_active_callbacks(client): ), f"{callback_name} not found in _active_callbacks={_active_callbacks}" assert not any( - "_ENTERPRISE_OpenAI_Moderation" in callback for callback in _active_callbacks - ), f"_ENTERPRISE_OpenAI_Moderation should not be in _active_callbacks={_active_callbacks}" + "ENTERPRISE_OpenAI_Moderation" in callback for callback in _active_callbacks + ), f"ENTERPRISE_OpenAI_Moderation should not be in _active_callbacks={_active_callbacks}" diff --git a/tests/unit/proxy/utils/proxy_logging/test_pre_call_hook.py b/tests/unit/proxy/utils/proxy_logging/test_pre_call_hook.py index e26eab0a759..2bf6a8d7e6c 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_pre_call_hook.py +++ b/tests/unit/proxy/utils/proxy_logging/test_pre_call_hook.py @@ -442,8 +442,8 @@ def test_every_pre_call_customlogger_is_deliberately_classified(): judges_content = { "_OPTIONAL_PromptInjectionDetection", "_PROXY_AzureContentSafety", - "_ENTERPRISE_BannedKeywords", - "_ENTERPRISE_BlockedUserList", + "ENTERPRISE_BannedKeywords", + "ENTERPRISE_BlockedUserList", } counts_or_shapes_the_request = { "_PROXY_MaxParallelRequestsHandler_v3", @@ -454,16 +454,16 @@ def test_every_pre_call_customlogger_is_deliberately_classified(): "_PROXY_SensitiveDataRoutingHandler", "ResponsesIDSecurity", "SkillsInjectionHook", - "_PROXY_LiteLLMManagedFiles", - "_PROXY_LiteLLMManagedVectorStores", + "PROXY_LiteLLMManagedFiles", + "PROXY_LiteLLMManagedVectorStores", } from litellm.proxy.hooks import PROXY_HOOKS registered = dict(PROXY_HOOKS) for name, cls in ( - ("banned_keywords", _load("enterprise.enterprise_hooks.banned_keywords", "_ENTERPRISE_BannedKeywords")), - ("blocked_user_check", _load("enterprise.enterprise_hooks.blocked_user_list", "_ENTERPRISE_BlockedUserList")), + ("banned_keywords", _load("enterprise.enterprise_hooks.banned_keywords", "ENTERPRISE_BannedKeywords")), + ("blocked_user_check", _load("enterprise.enterprise_hooks.blocked_user_list", "ENTERPRISE_BlockedUserList")), ("detect_prompt_injection", _load("litellm.proxy.hooks.prompt_injection_detection", "_OPTIONAL_PromptInjectionDetection")), ("azure_content_safety", _load("litellm.proxy.hooks.azure_content_safety", "_PROXY_AzureContentSafety")), ): diff --git a/tests/unit/proxy/vector_store_endpoints/test_vector_store_access_control.py b/tests/unit/proxy/vector_store_endpoints/test_vector_store_access_control.py index dc4f2900038..dabb89d6a46 100644 --- a/tests/unit/proxy/vector_store_endpoints/test_vector_store_access_control.py +++ b/tests/unit/proxy/vector_store_endpoints/test_vector_store_access_control.py @@ -161,7 +161,7 @@ async def _listed_ids(user_api_key_dict: UserAPIKeyAuth) -> list[str]: ) with patch( # test-quality-ok: the list route reads rows through this module-level DB helper, no injection seam - "litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry._get_vector_stores_from_db", + "litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry.get_vector_stores_from_db", new=AsyncMock(return_value=[_UNSCOPED, _TEAM_A_OWNED, _UI_CREATED]), ): response = await list_vector_stores(user_api_key_dict=user_api_key_dict) diff --git a/tests/unit/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/unit/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 2f7d4b350be..2f70c1f5188 100644 --- a/tests/unit/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/unit/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -1841,7 +1841,7 @@ async def test_vector_store_synchronization_across_instances(): # Step 3: Test that Instance 2 can list vector stores from database # (Simulate what happens in list_vector_stores endpoint - using DB as source of truth) - vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db( + vector_stores_from_db = await VectorStoreRegistry.get_vector_stores_from_db( prisma_client=mock_prisma_client ) @@ -1902,7 +1902,7 @@ async def test_vector_store_synchronization_across_instances(): # Step 5: Instance 2 should NOT show it in the list (database is source of truth) # The list endpoint logic should clean up stale cache entries vector_stores_from_db_after_delete = ( - await VectorStoreRegistry._get_vector_stores_from_db( + await VectorStoreRegistry.get_vector_stores_from_db( prisma_client=mock_prisma_client ) ) @@ -2054,7 +2054,7 @@ async def test_vector_store_update_and_list_synchronization(): instance_1_registry.add_vector_store_to_registry(vector_store=test_vector_store) # Step 2: Instance 2 fetches and caches the vector store - vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db( + vector_stores_from_db = await VectorStoreRegistry.get_vector_stores_from_db( prisma_client=mock_prisma_client ) for vs in vector_stores_from_db: @@ -2106,7 +2106,7 @@ async def test_vector_store_update_and_list_synchronization(): # Step 4: Instance 2 calls list endpoint (which should sync with database) # This simulates what list_vector_stores endpoint does vector_stores_from_db_after_update = ( - await VectorStoreRegistry._get_vector_stores_from_db( + await VectorStoreRegistry.get_vector_stores_from_db( prisma_client=mock_prisma_client ) ) diff --git a/tests/unit/proxy/vector_store_endpoints/test_vector_store_rbac.py b/tests/unit/proxy/vector_store_endpoints/test_vector_store_rbac.py index 87cfddd1ae3..f5934962f75 100644 --- a/tests/unit/proxy/vector_store_endpoints/test_vector_store_rbac.py +++ b/tests/unit/proxy/vector_store_endpoints/test_vector_store_rbac.py @@ -64,7 +64,7 @@ async def test_list_vector_stores_allowed_when_not_disabled(): with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with patch.object(litellm, "vector_store_registry", None): with patch( - "litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry._get_vector_stores_from_db", + "litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry.get_vector_stores_from_db", new=AsyncMock(return_value=[]), ): # Must not raise any HTTPException — if mocking is incomplete the @@ -119,7 +119,7 @@ async def test_list_vector_stores_admin_not_blocked(): with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with patch.object(litellm, "vector_store_registry", None): with patch( - "litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry._get_vector_stores_from_db", + "litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry.get_vector_stores_from_db", new=AsyncMock(return_value=[]), ): # Must not raise any HTTPException — admin is always allowed. @@ -161,7 +161,7 @@ async def test_list_vector_stores_accepts_non_positive_page_like_base(page): with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with patch.object(litellm, "vector_store_registry", None): with patch( - "litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry._get_vector_stores_from_db", + "litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry.get_vector_stores_from_db", new=AsyncMock(return_value=[]), ): response: Final = await list_vector_stores(user_api_key_dict=admin, page=page, page_size=10) diff --git a/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py index c13054f1cee..c17f5443025 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -1502,7 +1502,7 @@ class TestToolChoiceTransformation: def test_transform_tool_choice_for_responses_api_response( self, request_tool_choice: object, expected: str | dict[str, str] ) -> None: - result: Final = LiteLLMCompletionResponsesConfig._transform_tool_choice_for_responses_api_response( + result: Final = LiteLLMCompletionResponsesConfig.transform_tool_choice_for_responses_api_response( request_tool_choice ) assert result == expected @@ -2966,7 +2966,7 @@ class TestUsageTransformation: ) # Execute - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3008,7 +3008,7 @@ class TestUsageTransformation: ) # Execute - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3033,7 +3033,7 @@ class TestUsageTransformation: ), ) - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=usage ) @@ -3060,7 +3060,7 @@ class TestUsageTransformation: completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=10), ) - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=usage ) details = response_usage.input_tokens_details @@ -3071,7 +3071,7 @@ class TestUsageTransformation: from litellm.responses.utils import ResponseAPILoggingUtils - back = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response_usage.model_dump()) + back = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(response_usage.model_dump()) assert back.prompt_tokens_details.image_tokens == 150 assert back.prompt_tokens_details.video_tokens == 50 @@ -3104,7 +3104,7 @@ class TestUsageTransformation: ) # Execute - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3147,7 +3147,7 @@ class TestUsageTransformation: ) # Execute - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3195,7 +3195,7 @@ class TestUsageTransformation: ) # Execute - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3229,7 +3229,7 @@ class TestUsageTransformation: ) # Execute - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3272,7 +3272,7 @@ class TestUsageTransformation: ) # Execute - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3312,7 +3312,7 @@ class TestUsageTransformation: ], ) - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3348,7 +3348,7 @@ class TestUsageTransformation: ], ) - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3381,7 +3381,7 @@ class TestUsageTransformation: ], ) - response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + response_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( chat_completion_response=chat_completion_response ) @@ -3555,7 +3555,7 @@ class TestStreamingIDConsistency: mock_stream_wrapper = Mock(spec=litellm.CustomStreamWrapper) mock_logging_obj = Mock() mock_stream_wrapper.logging_obj = mock_logging_obj - mock_logging_obj._response_cost_calculator = Mock(return_value=0.001) + mock_logging_obj.response_cost_calculator = Mock(return_value=0.001) # Create the streaming iterator iterator = LiteLLMCompletionStreamingIterator( @@ -3764,7 +3764,7 @@ class TestCompletedResponseLatchedOnStreamEnd: mock_wrapper = Mock(spec=litellm.CustomStreamWrapper) mock_wrapper.logging_obj = Mock() - mock_wrapper.logging_obj._response_cost_calculator = Mock(return_value=0.0) + mock_wrapper.logging_obj.response_cost_calculator = Mock(return_value=0.0) mock_wrapper.__aiter__ = Mock(return_value=mock_wrapper) mock_wrapper.__anext__ = Mock(side_effect=StopAsyncIteration) diff --git a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py index 041bcf1b6d7..a42eb74b0fa 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py +++ b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py @@ -696,7 +696,7 @@ async def test_streaming_events_share_the_chat_completion_response_id(): response_ids = _response_ids(events) assert len(response_ids) == 3 assert len(set(response_ids)) == 1 - decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(response_ids[0]) + decoded = ResponsesAPIRequestUtils.decode_responses_api_response_id(response_ids[0]) assert decoded["response_id"] == CHAT_COMPLETION_ID assert decoded["custom_llm_provider"] == "anthropic" @@ -710,7 +710,7 @@ def test_sync_streaming_events_share_the_chat_completion_response_id(): assert len(response_ids) == 3 assert len(set(response_ids)) == 1 assert ( - ResponsesAPIRequestUtils._decode_responses_api_response_id(response_ids[0])["response_id"] + ResponsesAPIRequestUtils.decode_responses_api_response_id(response_ids[0])["response_id"] == CHAT_COMPLETION_ID ) diff --git a/tests/unit/responses/mcp/test_aresponses_api_with_mcp.py b/tests/unit/responses/mcp/test_aresponses_api_with_mcp.py index de0dc78af43..842b64ae380 100644 --- a/tests/unit/responses/mcp/test_aresponses_api_with_mcp.py +++ b/tests/unit/responses/mcp/test_aresponses_api_with_mcp.py @@ -55,15 +55,15 @@ async def test_mcp_helper_methods(): ] # Should return True for MCP tools with litellm_proxy - assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(mcp_tools) == True + assert LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(mcp_tools) == True # Should return False for other tools assert ( - LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(other_tools) == False + LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(other_tools) == False ) # Should return False for None - assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(None) == False + assert LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(None) == False # Test _parse_mcp_tools mixed_tools = mcp_tools + other_tools @@ -78,19 +78,19 @@ async def test_mcp_helper_methods(): mcp_tools_never = [{"require_approval": "never"}] mcp_tools_always = [{"require_approval": "always"}] - assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_never) == True + assert LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools(mcp_tools_never) == True assert ( - LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_always) == False + LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools(mcp_tools_always) == False ) # A single approval-required reference must disable auto-execution for the # whole request; otherwise a "never" reference alongside an "always" one # would let the approval-gated tool run without approval. mcp_tools_mixed = [{"require_approval": "never"}, {"require_approval": "always"}] - assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_mixed) == False + assert LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools(mcp_tools_mixed) == False mcp_tools_manual = [{"require_approval": "never"}, {"require_approval": "manual"}] - assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_manual) == False - assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools([]) == False + assert LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools(mcp_tools_manual) == False + assert LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools([]) == False print("✓ MCP helper methods test passed!") @@ -165,7 +165,7 @@ async def test_mcp_output_elements_addition(): ] # Test adding output elements - updated_response = LiteLLM_Proxy_MCP_Handler._add_mcp_output_elements_to_response( + updated_response = LiteLLM_Proxy_MCP_Handler.add_mcp_output_elements_to_response( response=mock_response, mcp_tools_fetched=mock_mcp_tools, tool_results=mock_tool_results, @@ -227,7 +227,7 @@ async def test_aresponses_api_with_mcp_mock_integration(): ) # Test 1: Verify MCP tools are detected correctly - should_use_mcp = LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway( + should_use_mcp = LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway( cast(Any, mcp_tools) ) assert ( @@ -235,7 +235,7 @@ async def test_aresponses_api_with_mcp_mock_integration(): ), "Should detect MCP tools with litellm_proxy server_url" # Test 2: Verify auto-execution detection works - should_auto_execute = LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools( + should_auto_execute = LiteLLM_Proxy_MCP_Handler.should_auto_execute_tools( cast(Any, mcp_tools) ) assert ( @@ -330,7 +330,7 @@ async def test_aresponses_api_with_mcp_passes_mcp_server_auth_headers_to_process with ( patch.object( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", mock_process, ), patch( @@ -733,7 +733,7 @@ async def test_streaming_mcp_events_validation(): ) as mock_get_tools, patch.object( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", new_callable=AsyncMock, ) as mock_execute_tools, patch( @@ -869,7 +869,7 @@ async def test_mcp_parameter_preparation_helpers(): } # Test Case 1: Auto-execute scenario (should disable streaming) - initial_params_auto = LiteLLM_Proxy_MCP_Handler._prepare_initial_call_params( + initial_params_auto = LiteLLM_Proxy_MCP_Handler.prepare_initial_call_params( call_params=base_call_params, should_auto_execute=True ) @@ -885,7 +885,7 @@ async def test_mcp_parameter_preparation_helpers(): print("✅ _prepare_initial_call_params (auto-execute) works correctly") # Test Case 2: No auto-execute scenario (should preserve streaming) - initial_params_no_auto = LiteLLM_Proxy_MCP_Handler._prepare_initial_call_params( + initial_params_no_auto = LiteLLM_Proxy_MCP_Handler.prepare_initial_call_params( call_params=base_call_params, should_auto_execute=False ) @@ -899,7 +899,7 @@ async def test_mcp_parameter_preparation_helpers(): print("✅ _prepare_initial_call_params (no auto-execute) works correctly") # Test _prepare_follow_up_call_params - follow_up_params = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params( + follow_up_params = LiteLLM_Proxy_MCP_Handler.prepare_follow_up_call_params( call_params=base_call_params, original_stream_setting=True ) @@ -1010,7 +1010,7 @@ async def test_mcp_tool_execution_events_creation(): ] # Create tool execution events - execution_events = LiteLLM_Proxy_MCP_Handler._create_tool_execution_events( + execution_events = LiteLLM_Proxy_MCP_Handler.create_tool_execution_events( tool_calls=mock_tool_calls, tool_results=mock_tool_results ) @@ -1036,7 +1036,7 @@ async def test_mcp_tool_execution_events_creation(): print("Tool execution events have proper structure") # Test with empty inputs - empty_events = LiteLLM_Proxy_MCP_Handler._create_tool_execution_events( + empty_events = LiteLLM_Proxy_MCP_Handler.create_tool_execution_events( tool_calls=[], tool_results=[] ) diff --git a/tests/unit/responses/mcp/test_chat_completions_handler.py b/tests/unit/responses/mcp/test_chat_completions_handler.py index 54e98cc2b6c..27477f24d92 100644 --- a/tests/unit/responses/mcp/test_chat_completions_handler.py +++ b/tests/unit/responses/mcp/test_chat_completions_handler.py @@ -44,7 +44,7 @@ async def test_acompletion_with_mcp_without_auto_execution_calls_model(monkeypat monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_use_litellm_mcp_gateway", + "should_use_litellm_mcp_gateway", staticmethod(lambda tools: True), ) monkeypatch.setattr( @@ -58,17 +58,17 @@ async def test_acompletion_with_mcp_without_auto_execution_calls_model(monkeypat monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", mock_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_transform_mcp_tools_to_openai", + "transform_mcp_tools_to_openai", staticmethod(lambda *_, **__: ["openai-tool"]), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_auto_execute_tools", + "should_auto_execute_tools", staticmethod(lambda **_: False), ) captured_secret_fields = {} @@ -119,7 +119,7 @@ async def test_acompletion_with_mcp_passes_mcp_server_auth_headers_to_process_to monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_use_litellm_mcp_gateway", + "should_use_litellm_mcp_gateway", staticmethod(lambda t: True), ) monkeypatch.setattr( @@ -129,17 +129,17 @@ async def test_acompletion_with_mcp_passes_mcp_server_auth_headers_to_process_to ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", mock_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_transform_mcp_tools_to_openai", + "transform_mcp_tools_to_openai", staticmethod(lambda *_, **__: ["openai-tool"]), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_auto_execute_tools", + "should_auto_execute_tools", staticmethod(lambda **_: False), ) @@ -292,7 +292,7 @@ async def test_acompletion_with_mcp_auto_exec_performs_follow_up(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_use_litellm_mcp_gateway", + "should_use_litellm_mcp_gateway", staticmethod(lambda tools: True), ) monkeypatch.setattr( @@ -306,22 +306,22 @@ async def test_acompletion_with_mcp_auto_exec_performs_follow_up(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", mock_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_transform_mcp_tools_to_openai", + "transform_mcp_tools_to_openai", staticmethod(lambda *_, **__: tools), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_auto_execute_tools", + "should_auto_execute_tools", staticmethod(lambda **_: True), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_extract_tool_calls_from_chat_response", + "extract_tool_calls_from_chat_response", staticmethod( lambda **_: [ { @@ -338,12 +338,12 @@ async def test_acompletion_with_mcp_auto_exec_performs_follow_up(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", mock_execute, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_create_follow_up_messages_for_chat", + "create_follow_up_messages_for_chat", staticmethod( lambda **_: [ {"role": "user", "content": "hello"}, @@ -499,7 +499,7 @@ async def test_acompletion_with_mcp_adds_metadata_to_streaming(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_use_litellm_mcp_gateway", + "should_use_litellm_mcp_gateway", staticmethod(lambda tools: True), ) monkeypatch.setattr( @@ -513,17 +513,17 @@ async def test_acompletion_with_mcp_adds_metadata_to_streaming(monkeypatch): monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", mock_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_transform_mcp_tools_to_openai", + "transform_mcp_tools_to_openai", staticmethod(lambda *_, **__: openai_tools), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_auto_execute_tools", + "should_auto_execute_tools", staticmethod(lambda **_: False), ) monkeypatch.setattr( @@ -638,7 +638,7 @@ async def test_acompletion_with_mcp_streaming_initial_call_is_streaming(monkeypa monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_use_litellm_mcp_gateway", + "should_use_litellm_mcp_gateway", staticmethod(lambda tools: True), ) monkeypatch.setattr( @@ -652,22 +652,22 @@ async def test_acompletion_with_mcp_streaming_initial_call_is_streaming(monkeypa monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", mock_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_transform_mcp_tools_to_openai", + "transform_mcp_tools_to_openai", staticmethod(lambda *_, **__: openai_tools), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_auto_execute_tools", + "should_auto_execute_tools", staticmethod(lambda **_: True), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_extract_tool_calls_from_chat_response", + "extract_tool_calls_from_chat_response", staticmethod( lambda **_: [ { @@ -684,12 +684,12 @@ async def test_acompletion_with_mcp_streaming_initial_call_is_streaming(monkeypa monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", mock_execute, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_create_follow_up_messages_for_chat", + "create_follow_up_messages_for_chat", staticmethod( lambda **_: [ {"role": "user", "content": "hello"}, @@ -880,7 +880,7 @@ async def test_acompletion_with_mcp_streaming_metadata_in_correct_chunks(monkeyp monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_use_litellm_mcp_gateway", + "should_use_litellm_mcp_gateway", staticmethod(lambda tools: True), ) monkeypatch.setattr( @@ -894,22 +894,22 @@ async def test_acompletion_with_mcp_streaming_metadata_in_correct_chunks(monkeyp monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", mock_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_transform_mcp_tools_to_openai", + "transform_mcp_tools_to_openai", staticmethod(lambda *_, **__: openai_tools), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_auto_execute_tools", + "should_auto_execute_tools", staticmethod(lambda **_: True), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_extract_tool_calls_from_chat_response", + "extract_tool_calls_from_chat_response", staticmethod(lambda **_: tool_calls), ) @@ -918,12 +918,12 @@ async def test_acompletion_with_mcp_streaming_metadata_in_correct_chunks(monkeyp monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_execute_tool_calls", + "execute_tool_calls", mock_execute, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_create_follow_up_messages_for_chat", + "create_follow_up_messages_for_chat", staticmethod( lambda **_: [ {"role": "user", "content": "hello"}, @@ -1094,7 +1094,7 @@ async def test_execute_tool_calls_sets_proxy_server_request_arguments(monkeypatc user_api_key_auth.api_key = "test_key" # Call _execute_tool_calls - result = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + result = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map=tool_server_map, tool_calls=tool_calls, user_api_key_auth=user_api_key_auth, @@ -1182,7 +1182,7 @@ async def test_acompletion_with_mcp_streaming_drain_error_does_not_drop_final_ch monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_use_litellm_mcp_gateway", + "should_use_litellm_mcp_gateway", staticmethod(lambda tools: True), ) monkeypatch.setattr( @@ -1196,22 +1196,22 @@ async def test_acompletion_with_mcp_streaming_drain_error_does_not_drop_final_ch monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", mock_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_transform_mcp_tools_to_openai", + "transform_mcp_tools_to_openai", staticmethod(lambda *_, **__: openai_tools), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_auto_execute_tools", + "should_auto_execute_tools", staticmethod(lambda **_: True), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_extract_tool_calls_from_chat_response", + "extract_tool_calls_from_chat_response", staticmethod(lambda **_: []), ) monkeypatch.setattr( @@ -1300,7 +1300,7 @@ async def test_acompletion_with_mcp_streaming_drains_inner_stream_after_exhausti monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_use_litellm_mcp_gateway", + "should_use_litellm_mcp_gateway", staticmethod(lambda tools: True), ) monkeypatch.setattr( @@ -1314,22 +1314,22 @@ async def test_acompletion_with_mcp_streaming_drains_inner_stream_after_exhausti monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", mock_process, ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_transform_mcp_tools_to_openai", + "transform_mcp_tools_to_openai", staticmethod(lambda *_, **__: openai_tools), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_should_auto_execute_tools", + "should_auto_execute_tools", staticmethod(lambda **_: True), ) monkeypatch.setattr( LiteLLM_Proxy_MCP_Handler, - "_extract_tool_calls_from_chat_response", + "extract_tool_calls_from_chat_response", staticmethod(lambda **_: []), ) monkeypatch.setattr( diff --git a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index 73bee304fc2..7445df1b066 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -123,7 +123,7 @@ def test_extract_tool_calls_from_chat_response_handles_tool_calls(): object="chat.completion", ) - tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(response) + tool_calls = LiteLLM_Proxy_MCP_Handler.extract_tool_calls_from_chat_response(response) assert len(tool_calls) == 1 assert tool_calls[0]["function"]["name"] == "foo" @@ -161,7 +161,7 @@ def test_create_follow_up_messages_for_chat_appends_tool_results(): } ] - follow_up = LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( + follow_up = LiteLLM_Proxy_MCP_Handler.create_follow_up_messages_for_chat( original_messages, response, tool_results, @@ -193,8 +193,8 @@ def test_transform_mcp_tools_to_openai_uses_chat_format(monkeypatch): fake_transform_responses, ) - chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"], target_format="chat") - resp_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"]) + chat_tools = LiteLLM_Proxy_MCP_Handler.transform_mcp_tools_to_openai(["tool"], target_format="chat") + resp_tools = LiteLLM_Proxy_MCP_Handler.transform_mcp_tools_to_openai(["tool"]) assert chat_tools == [{"chat": True}] assert resp_tools == [{"responses": True}] @@ -216,7 +216,7 @@ def test_create_follow_up_input_handles_response_function_tool_call(): ] ) - follow_up = LiteLLM_Proxy_MCP_Handler._create_follow_up_input( + follow_up = LiteLLM_Proxy_MCP_Handler.create_follow_up_input( response=cast(Any, response), tool_results=[], original_input=None, @@ -243,7 +243,7 @@ async def test_execute_tool_calls_strips_server_prefix(monkeypatch): } ] - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=tool_calls, user_api_key_auth=None, @@ -265,7 +265,7 @@ async def test_execute_tool_calls_keeps_tool_name_without_prefix(monkeypatch): } ] - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=tool_calls, user_api_key_auth=None, @@ -287,7 +287,7 @@ async def test_execute_tool_calls_keeps_tool_name_when_equal_to_server(monkeypat } ] - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "echo"}, tool_calls=tool_calls, user_api_key_auth=None, @@ -323,7 +323,7 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n } ] - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki_test"}, tool_calls=tool_calls, user_api_key_auth=None, @@ -366,7 +366,7 @@ async def test_execute_tool_calls_reverse_maps_display_name(monkeypatch): } ] - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki_mcp"}, tool_calls=tool_calls, user_api_key_auth=None, @@ -399,7 +399,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") - results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + results = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=tool_calls, user_api_key_auth=user_auth, @@ -442,7 +442,7 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio tool_name = "deepwiki-read_wiki_structure" tool_calls = [{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}] - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=tool_calls, user_api_key_auth=None, @@ -482,7 +482,7 @@ async def test_execute_tool_calls_threads_logging_obj_into_call_tool(monkeypatch tool_name = "deepwiki-read_wiki_structure" tool_calls = [{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}] - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=tool_calls, user_api_key_auth=None, @@ -524,7 +524,7 @@ async def test_execute_tool_calls_applies_post_call_hook_content(monkeypatch): monkeypatch.setattr(handler_module, "function_setup", lambda *_args, **_kwargs: (logging_obj, None)) tool_name = "deepwiki-read_wiki_structure" - results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + results = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=[{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}], user_api_key_auth=None, @@ -558,7 +558,7 @@ async def test_execute_tool_calls_returns_proxy_result_without_logging(monkeypat monkeypatch.setattr(handler_module, "function_setup", lambda *_args, **_kwargs: (None, None)) tool_name = "deepwiki-read_wiki_structure" - results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + results = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=[{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}], user_api_key_auth=None, @@ -595,7 +595,7 @@ async def test_execute_tool_calls_passes_logging_details_to_proxy_hook(monkeypat monkeypatch.setattr(handler_module, "function_setup", lambda *_args, **_kwargs: (logging_obj, None)) tool_name = "deepwiki-read_wiki_structure" - results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + results = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=[{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}], user_api_key_auth=None, @@ -637,7 +637,7 @@ async def test_execute_tool_calls_continues_when_post_call_logging_fails(monkeyp monkeypatch.setattr(handler_module, "function_setup", lambda *_args, **_kwargs: (logging_obj, None)) tool_name = "deepwiki-read_wiki_structure" - results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + results = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=[{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}], user_api_key_auth=None, @@ -684,7 +684,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}], ) - forwarded: Final = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(tools) + forwarded: Final = LiteLLM_Proxy_MCP_Handler.transform_mcp_tools_to_openai(tools) assert [tool["name"] for tool in forwarded] == ["safe", "masked"] assert forwarded[0]["description"] == "Safe lookup" assert forwarded[1]["description"] == "Contact [MASKED]" @@ -700,12 +700,12 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch def test_get_parent_request_tags_from_metadata(): - tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags({"metadata": {"tags": ["team-a", "prod"]}}) + tags = LiteLLM_Proxy_MCP_Handler.get_parent_request_tags({"metadata": {"tags": ["team-a", "prod"]}}) assert tags == ["team-a", "prod"] def test_get_parent_request_tags_from_nested_litellm_params(): - tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags( + tags = LiteLLM_Proxy_MCP_Handler.get_parent_request_tags( { "metadata": {"tags": ["top-level"]}, "litellm_params": { @@ -761,7 +761,7 @@ async def test_execute_tool_calls_exposes_sanitized_client_headers_to_logging(mo monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=[{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}], user_api_key_auth=None, @@ -787,7 +787,7 @@ async def test_execute_tool_calls_propagates_request_tags_to_function_setup(monk monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" - await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, tool_calls=[{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}], user_api_key_auth=None, @@ -863,7 +863,7 @@ def test_extract_tool_call_details_reads_anthropic_tool_use_input(): "input": {"repoName": "BerriAI/litellm"}, } - name, arguments, call_id = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_use_block) + name, arguments, call_id = LiteLLM_Proxy_MCP_Handler.extract_tool_call_details(tool_use_block) assert name == "read_wiki_structure" assert call_id == "toolu_01ABC" @@ -878,7 +878,7 @@ def test_extract_tool_call_details_still_prefers_openai_arguments(): "function": {"name": "get_weather", "arguments": '{"city": "Paris"}'}, } - name, arguments, call_id = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(openai_tool_call) + name, arguments, call_id = LiteLLM_Proxy_MCP_Handler.extract_tool_call_details(openai_tool_call) assert name == "get_weather" assert call_id == "call_123" @@ -978,7 +978,7 @@ async def test_split_mcp_tools_leaves_external_mcp_path_urls_for_the_provider(): assert names == {"mcp"} return frozenset() - gateway_tools, other_tools = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools( + gateway_tools, other_tools = await LiteLLM_Proxy_MCP_Handler.split_mcp_tools( [ZAPIER_TOOL, EXPLICIT_GATEWAY_TOOL, FUNCTION_TOOL], served_names=served_names ) @@ -1000,7 +1000,7 @@ async def test_split_mcp_tools_repoints_served_proxy_urls_at_the_gateway(): async def served_names(names): return frozenset({"my-toolset"}) - gateway_tools, other_tools = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools( + gateway_tools, other_tools = await LiteLLM_Proxy_MCP_Handler.split_mcp_tools( [served_tool, unserved_tool], served_names=served_names ) @@ -1013,7 +1013,7 @@ async def test_split_mcp_tools_skips_resolution_when_nothing_points_at_the_proxy async def served_names(names): raise AssertionError("no lookup expected") - gateway_tools, other_tools = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools( + gateway_tools, other_tools = await LiteLLM_Proxy_MCP_Handler.split_mcp_tools( [EXPLICIT_GATEWAY_TOOL, FUNCTION_TOOL], served_names=served_names ) @@ -1022,10 +1022,10 @@ async def test_split_mcp_tools_skips_resolution_when_nothing_points_at_the_proxy def test_should_use_gateway_still_triggers_on_http_mcp_path(): - assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway([ZAPIER_TOOL]) is True - assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway([EXPLICIT_GATEWAY_TOOL]) is True - assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway([FUNCTION_TOOL]) is False - assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(None) is False + assert LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway([ZAPIER_TOOL]) is True + assert LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway([EXPLICIT_GATEWAY_TOOL]) is True + assert LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway([FUNCTION_TOOL]) is False + assert LiteLLM_Proxy_MCP_Handler.should_use_litellm_mcp_gateway(None) is False @pytest.mark.asyncio @@ -1092,7 +1092,7 @@ def test_create_follow_up_input_preserves_reasoning_when_stateless(): Regression test (LIT-5427): a store=false follow-up has to replay the reasoning item, including reasoning.encrypted_content, since the provider kept no state. """ - follow_up = LiteLLM_Proxy_MCP_Handler._create_follow_up_input( + follow_up = LiteLLM_Proxy_MCP_Handler.create_follow_up_input( response=_response_with_reasoning_and_tool_call(), tool_results=[{"tool_call_id": "call-1", "name": "foo", "result": "done"}], original_input="hi", @@ -1144,7 +1144,7 @@ def test_create_follow_up_input_keeps_each_reasoning_item_before_its_function_ca item that follows it, so the replay has to keep the response's output order instead of grouping every reasoning item ahead of every function call. """ - follow_up = LiteLLM_Proxy_MCP_Handler._create_follow_up_input( + follow_up = LiteLLM_Proxy_MCP_Handler.create_follow_up_input( response=_response_with_interleaved_reasoning_and_tool_calls(), tool_results=[ {"tool_call_id": "call-1", "name": "foo", "result": "one"}, @@ -1175,7 +1175,7 @@ def test_create_follow_up_input_keeps_each_reasoning_item_before_its_function_ca def test_create_follow_up_input_omits_reasoning_when_stateful(): """With store=true the provider still holds the reasoning item, so don't resend it.""" - follow_up = LiteLLM_Proxy_MCP_Handler._create_follow_up_input( + follow_up = LiteLLM_Proxy_MCP_Handler.create_follow_up_input( response=_response_with_reasoning_and_tool_call(), tool_results=[{"tool_call_id": "call-1", "name": "foo", "result": "done"}], original_input="hi", @@ -1194,7 +1194,7 @@ def test_create_follow_up_input_omits_reasoning_when_stateful(): ], ) def test_is_persistence_disabled(call_params: dict[str, Any], expected: bool): - assert LiteLLM_Proxy_MCP_Handler._is_persistence_disabled(call_params) is expected + assert LiteLLM_Proxy_MCP_Handler.is_persistence_disabled(call_params) is expected @pytest.mark.parametrize( @@ -1251,9 +1251,9 @@ async def test_mcp_follow_up_call_is_stateless_when_store_is_false( monkeypatch.setattr(responses_main, "aresponses", fake_aresponses) monkeypatch.setattr(mcp_handler_module, "aresponses", fake_aresponses) monkeypatch.setattr( - LiteLLM_Proxy_MCP_Handler, "_process_mcp_tools_without_openai_transform", staticmethod(fake_process) + LiteLLM_Proxy_MCP_Handler, "process_mcp_tools_without_openai_transform", staticmethod(fake_process) ) - monkeypatch.setattr(LiteLLM_Proxy_MCP_Handler, "_execute_tool_calls", staticmethod(fake_execute)) + monkeypatch.setattr(LiteLLM_Proxy_MCP_Handler, "execute_tool_calls", staticmethod(fake_execute)) await responses_main.aresponses_api_with_mcp( input="hi", @@ -1372,7 +1372,7 @@ async def test_bridge_listing_leaves_the_callers_catalog_unchanged( if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None } assert bool(before) is real_listing - tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + tools, _server_names = await LiteLLM_Proxy_MCP_Handler.process_mcp_tools_without_openai_transform( user_api_key_auth=user, mcp_tools_with_litellm_proxy=[ { @@ -1448,7 +1448,7 @@ async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch async def bridge(name: str, first: bool) -> None: if not first: await first_listed.wait() - tools, server_map = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + tools, server_map = await LiteLLM_Proxy_MCP_Handler.process_mcp_tools_without_openai_transform( user_api_key_auth=user, mcp_tools_with_litellm_proxy=[ {"type": "mcp", "server_url": "litellm_proxy/mcp/bridge", "allowed_tools": [name]} @@ -1456,7 +1456,7 @@ async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch ) (first_listed if first else second_listed).set() await second_listed.wait() - result: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + result: Final = await LiteLLM_Proxy_MCP_Handler.execute_tool_calls( tool_server_map=server_map, tool_calls=[ {"type": "function_call", "name": f"bridge-{name}", "arguments": '{"probe":"bridge"}', "call_id": name} diff --git a/tests/unit/responses/test_responses_prompt_management.py b/tests/unit/responses/test_responses_prompt_management.py index 4379f4f28d3..60376408326 100644 --- a/tests/unit/responses/test_responses_prompt_management.py +++ b/tests/unit/responses/test_responses_prompt_management.py @@ -71,7 +71,7 @@ def _patch_responses_dispatch(): side_effect=_provider_by_model, ), patch.object( - import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, "_should_use_litellm_mcp_gateway", + import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, "should_use_litellm_mcp_gateway", return_value=False, ), patch.object( diff --git a/tests/unit/responses/test_responses_router_cooldown.py b/tests/unit/responses/test_responses_router_cooldown.py index e173c174521..e4bad083c87 100644 --- a/tests/unit/responses/test_responses_router_cooldown.py +++ b/tests/unit/responses/test_responses_router_cooldown.py @@ -13,7 +13,7 @@ import pytest import litellm -from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments +from litellm.router_utils.cooldown_handlers import async_get_cooldown_deployments @pytest.mark.asyncio @@ -76,7 +76,7 @@ async def test_responses_api_rate_limit_marks_deployment_for_cooldown(): input="hi", ) - cooldown_ids = await _async_get_cooldown_deployments( + cooldown_ids = await async_get_cooldown_deployments( litellm_router_instance=router, parent_otel_span=None ) assert failing_deployment_id in cooldown_ids, ( diff --git a/tests/unit/responses/test_responses_utils.py b/tests/unit/responses/test_responses_utils.py index 7a482488706..55d90acaaee 100644 --- a/tests/unit/responses/test_responses_utils.py +++ b/tests/unit/responses/test_responses_utils.py @@ -205,13 +205,13 @@ class TestResponsesAPIRequestUtils: """Ensure _update_responses_api_response_id_with_model_id works with dict input""" responses_api_response = {"id": "resp_abc123"} litellm_metadata = {"model_info": {"id": "gpt-4o"}} - updated = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + updated = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( responses_api_response=responses_api_response, custom_llm_provider="openai", litellm_metadata=litellm_metadata, ) assert updated["id"] != "resp_abc123" - decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(updated["id"]) + decoded = ResponsesAPIRequestUtils.decode_responses_api_response_id(updated["id"]) assert decoded.get("response_id") == "resp_abc123" assert decoded.get("model_id") == "gpt-4o" assert decoded.get("custom_llm_provider") == "openai" @@ -220,12 +220,12 @@ class TestResponsesAPIRequestUtils: raw = "resp_" + "a" * 48 litellm_metadata = {"model_info": {"id": "model-123"}} - once = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + once = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( {"id": raw}, custom_llm_provider="openai", litellm_metadata=litellm_metadata, ) - twice = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + twice = ResponsesAPIRequestUtils.update_responses_api_response_id_with_model_id( {"id": once["id"]}, custom_llm_provider="openai", litellm_metadata=litellm_metadata, @@ -233,17 +233,17 @@ class TestResponsesAPIRequestUtils: assert twice == once assert ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(twice["id"]) == raw - assert ResponsesAPIRequestUtils._decode_responses_api_response_id(once["id"]).get("response_id") == raw + assert ResponsesAPIRequestUtils.decode_responses_api_response_id(once["id"]).get("response_id") == raw def test_build_decode_container_id_omits_none_model_id(self): """model_id=None must not round-trip as the truthy string 'None'.""" - encoded = ResponsesAPIRequestUtils._build_container_id( + encoded = ResponsesAPIRequestUtils.build_container_id( custom_llm_provider="azure", model_id=None, container_id="cntr_upstream_abc", ) assert "None" not in base64.b64decode(encoded.replace("cntr_", "").encode("utf-8")).decode("utf-8") - decoded = ResponsesAPIRequestUtils._decode_container_id(encoded) + decoded = ResponsesAPIRequestUtils.decode_container_id(encoded) assert decoded.get("custom_llm_provider") == "azure" assert decoded.get("model_id") is None assert decoded.get("response_id") == "cntr_upstream_abc" @@ -252,7 +252,7 @@ class TestResponsesAPIRequestUtils: """IDs encoded before the None fix should decode without a bogus model_id.""" legacy_inner = "litellm:custom_llm_provider:azure;model_id:None;container_id:cntr_x" legacy_id = "cntr_" + base64.b64encode(legacy_inner.encode("utf-8")).decode("utf-8") - decoded = ResponsesAPIRequestUtils._decode_container_id(legacy_id) + decoded = ResponsesAPIRequestUtils.decode_container_id(legacy_id) assert decoded.get("model_id") is None assert decoded.get("custom_llm_provider") == "azure" assert decoded.get("response_id") == "cntr_x" @@ -265,7 +265,7 @@ class TestResponseAPILoggingUtils: usage = {"input_tokens": 10, "output_tokens": 20} # Execute - result = ResponseAPILoggingUtils._is_response_api_usage(usage) + result = ResponseAPILoggingUtils.is_response_api_usage(usage) # Assert assert result is True @@ -276,7 +276,7 @@ class TestResponseAPILoggingUtils: usage = {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30} # Execute - result = ResponseAPILoggingUtils._is_response_api_usage(usage) + result = ResponseAPILoggingUtils.is_response_api_usage(usage) # Assert assert result is False @@ -293,7 +293,7 @@ class TestResponseAPILoggingUtils: } # Execute - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) # Assert assert isinstance(result, Usage) @@ -313,7 +313,7 @@ class TestResponseAPILoggingUtils: } # Execute - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) # Assert assert result.prompt_tokens == 0 @@ -332,7 +332,7 @@ class TestResponseAPILoggingUtils: } # Execute - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) # Assert assert result.prompt_tokens == 15 @@ -369,7 +369,7 @@ class TestResponseAPILoggingUtils: } # Execute - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) # Assert - verify basic token counts assert isinstance(result, Usage) @@ -404,7 +404,7 @@ class TestResponseAPILoggingUtils: }, } - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result.prompt_tokens_details is not None assert result.prompt_tokens_details.cache_write_tokens == 10059 @@ -433,7 +433,7 @@ class TestResponseAPILoggingUtils: } # Execute - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) # Assert - all token detail types should be preserved assert result.prompt_tokens_details is not None @@ -465,7 +465,7 @@ class TestResponseAPILoggingUtils: }, } - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result.prompt_tokens_details is not None assert result.prompt_tokens_details.text_tokens == 8 @@ -487,7 +487,7 @@ class TestResponseAPILoggingUtils: "output_token_details": {"text_tokens": 2, "audio_tokens": 98}, } - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result.prompt_tokens_details is not None assert result.prompt_tokens_details.text_tokens == 10 @@ -507,7 +507,7 @@ class TestResponseAPILoggingUtils: "output_token_details": {"text_tokens": 70, "audio_tokens": 0, "reasoning_tokens": 52}, } - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result.completion_tokens == 70 assert result.completion_tokens_details is not None @@ -525,7 +525,7 @@ class TestResponseAPILoggingUtils: "output_token_details": {"text_tokens": 39, "audio_tokens": 31, "reasoning_tokens": 23}, } - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result.completion_tokens_details is not None assert result.completion_tokens_details.text_tokens == 16 @@ -541,7 +541,7 @@ class TestResponseAPILoggingUtils: "output_tokens_details": {"text_tokens": 12, "reasoning_tokens": 5}, } - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result.completion_tokens_details is not None assert result.completion_tokens_details.text_tokens == 12 @@ -557,7 +557,7 @@ class TestResponseAPILoggingUtils: server_side_tool_usage_details=details, ) - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert isinstance(result, Usage) assert result.prompt_tokens == 100 @@ -577,7 +577,7 @@ class TestResponseAPILoggingUtils: server_side_tool_usage_details={"web_search_calls": 1}, ) - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result.prompt_tokens == 35 assert result.completion_tokens == 1716 @@ -594,7 +594,7 @@ class TestResponseAPILoggingUtils: ) setattr(usage, "server_side_tool_usage_details", details) - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result is usage assert getattr(result, "server_side_tool_usage_details") == details @@ -621,7 +621,7 @@ class TestResponseAPILoggingUtils: "server_side_tool_usage_details": details, } - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert isinstance(result, Usage) assert result.prompt_tokens == 50 @@ -647,7 +647,7 @@ class TestResponseAPILoggingUtils: }, } - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result.prompt_tokens_details is not None assert result.prompt_tokens_details.cached_tokens == 192 @@ -668,7 +668,7 @@ class TestResponseAPILoggingUtils: }, ) - result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage) + result = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(usage) assert result.prompt_tokens_details is not None assert result.prompt_tokens_details.cached_tokens_details is not None diff --git a/tests/unit/responses/test_responses_websocket_all_providers.py b/tests/unit/responses/test_responses_websocket_all_providers.py index 222cc641b6b..b29e9a6b7db 100644 --- a/tests/unit/responses/test_responses_websocket_all_providers.py +++ b/tests/unit/responses/test_responses_websocket_all_providers.py @@ -2974,7 +2974,7 @@ def _wrapped_reasoning_item(): return { "type": "reasoning", "id": ResponsesAPIRequestUtils._build_encrypted_item_id("dep-1", "rs_orig"), - "encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAA-blob", "dep-1"), + "encrypted_content": ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAA-blob", "dep-1"), "summary": [], } @@ -3063,7 +3063,7 @@ class TestNativeWebSocketEncryptedContentAffinity: await handler.backend_to_client() - wrapped_content = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAA-blob", "dep-1") + wrapped_content = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAA-blob", "dep-1") item_done = json.loads(websocket.send_text.await_args_list[0][0][0]) assert item_done["item"]["encrypted_content"] == wrapped_content completed = json.loads(websocket.send_text.await_args_list[1][0][0]) @@ -3166,7 +3166,7 @@ class TestNativeWebSocketEncryptedContentAffinity: logging_obj = MagicMock() logging_obj.dispatch_success_handlers = AsyncMock() logging_obj.dispatch_failure_handlers = AsyncMock() - logging_obj._response_cost_calculator = MagicMock(return_value=0.0) + logging_obj.response_cost_calculator = MagicMock(return_value=0.0) handler = _make_streaming( websocket=websocket, backend_ws=backend_ws, @@ -3215,7 +3215,7 @@ class TestNativeWebSocketEncryptedContentAffinity: logging_obj = MagicMock() logging_obj.dispatch_success_handlers = AsyncMock() logging_obj.dispatch_failure_handlers = AsyncMock() - logging_obj._response_cost_calculator = MagicMock(return_value=0.01) + logging_obj.response_cost_calculator = MagicMock(return_value=0.01) handler = _make_streaming(websocket=websocket, backend_ws=backend_ws, logging_obj=logging_obj, request_data={}) await handler.backend_to_client() @@ -3271,7 +3271,7 @@ class TestNativeWebSocketEncryptedContentAffinity: logging_obj = MagicMock() logging_obj.dispatch_success_handlers = AsyncMock() logging_obj.dispatch_failure_handlers = AsyncMock() - logging_obj._response_cost_calculator = MagicMock(return_value=0.0) + logging_obj.response_cost_calculator = MagicMock(return_value=0.0) handler = _make_streaming( websocket=websocket, backend_ws=backend_ws, diff --git a/tests/unit/responses/test_streaming_iterator.py b/tests/unit/responses/test_streaming_iterator.py index 9e9d221ce84..38ed68a9abc 100644 --- a/tests/unit/responses/test_streaming_iterator.py +++ b/tests/unit/responses/test_streaming_iterator.py @@ -140,7 +140,7 @@ async def test_responses_streaming_stamps_completion_start_time_on_first_chunk() logging_obj.completion_start_time = completion_start_time logging_obj.model_call_details["completion_start_time"] = completion_start_time - logging_obj._update_completion_start_time.side_effect = _update + logging_obj.update_completion_start_time.side_effect = _update iterator = _make_iterator( sse_events=[ @@ -182,7 +182,7 @@ async def test_responses_streaming_does_not_reset_prior_completion_start_time(): async for _ in iterator: pass - logging_obj._update_completion_start_time.assert_not_called() + logging_obj.update_completion_start_time.assert_not_called() assert logging_obj.completion_start_time == prior @@ -628,7 +628,7 @@ def test_stream_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypat ) logging_obj = SimpleNamespace( model_call_details={"litellm_params": {}}, - _llm_caching_handler=caching_handler, + llm_caching_handler=caching_handler, ) iterator = ResponsesAPIStreamingIterator( response=httpx.Response(200), @@ -787,12 +787,12 @@ def test_stamp_responses_usage_cost_stamps_computed_cost(): response = _responses_api_response_with_usage() logging_obj = Mock(spec=LiteLLMLoggingObj) - logging_obj._response_cost_calculator.return_value = 0.000704 + logging_obj.response_cost_calculator.return_value = 0.000704 _stamp_responses_usage_cost(response, logging_obj) assert getattr(response.usage, "cost", None) == pytest.approx(0.000704) - logging_obj._response_cost_calculator.assert_called_once_with(result=response) + logging_obj.response_cost_calculator.assert_called_once_with(result=response) def test_stamp_responses_usage_cost_keeps_provider_reported_cost(): @@ -805,7 +805,7 @@ def test_stamp_responses_usage_cost_keeps_provider_reported_cost(): _stamp_responses_usage_cost(response, logging_obj) assert getattr(response.usage, "cost", None) == pytest.approx(0.5) - logging_obj._response_cost_calculator.assert_not_called() + logging_obj.response_cost_calculator.assert_not_called() def _unvalidated_response_with_dict_usage(usage: dict) -> ResponsesAPIResponse: @@ -839,20 +839,20 @@ def test_stamp_responses_usage_cost_keeps_provider_cost_from_dict_usage(): assert isinstance(response.usage, ResponseAPIUsage) assert response.usage.cost == pytest.approx(3e-05) assert response.usage.output_tokens_details.reasoning_tokens == 117 - logging_obj._response_cost_calculator.assert_not_called() + logging_obj.response_cost_calculator.assert_not_called() def test_stamp_responses_usage_cost_computes_cost_for_dict_usage_without_cost(): from litellm.responses.streaming_iterator import _stamp_responses_usage_cost response = _unvalidated_response_with_dict_usage({"input_tokens": 29, "output_tokens": 120, "total_tokens": 149}) logging_obj = Mock(spec=LiteLLMLoggingObj) - logging_obj._response_cost_calculator.return_value = 0.000704 + logging_obj.response_cost_calculator.return_value = 0.000704 _stamp_responses_usage_cost(response, logging_obj) assert isinstance(response.usage, ResponseAPIUsage) assert response.usage.cost == pytest.approx(0.000704) - logging_obj._response_cost_calculator.assert_called_once_with(result=response) + logging_obj.response_cost_calculator.assert_called_once_with(result=response) def test_stamp_responses_usage_cost_survives_calculator_failure(): @@ -860,7 +860,7 @@ def test_stamp_responses_usage_cost_survives_calculator_failure(): response = _responses_api_response_with_usage() logging_obj = Mock(spec=LiteLLMLoggingObj) - logging_obj._response_cost_calculator.side_effect = RuntimeError("cost map unavailable") + logging_obj.response_cost_calculator.side_effect = RuntimeError("cost map unavailable") _stamp_responses_usage_cost(response, logging_obj) @@ -1223,7 +1223,7 @@ async def test_completed_event_with_a_dict_response_is_typed_and_billed(): config: Final = Mock(spec=BaseResponsesAPIConfig) config.transform_streaming_response.side_effect = _transform logging_obj: Final = _logging_obj_stub() - logging_obj._response_cost_calculator.return_value = 0.000704 + logging_obj.response_cost_calculator.return_value = 0.000704 iterator: Final = _make_iterator( sse_events=[ _sse_event({"type": "response.output_text.delta", "delta": "hello world"}), @@ -1245,7 +1245,7 @@ async def test_completed_event_with_a_dict_response_is_typed_and_billed(): assert usage.input_tokens > 0 assert usage.output_tokens > 0 assert usage.cost == pytest.approx(0.000704) - logging_obj._response_cost_calculator.assert_any_call(result=completed_response) + logging_obj.response_cost_calculator.assert_any_call(result=completed_response) def test_billed_terminal_response_keeps_a_response_that_already_has_usage(): @@ -1273,7 +1273,7 @@ def test_persist_completed_response_to_cache_skips_a_response_without_output(mon logging_obj: Final = _logging_obj_stub() caching_handler: Final = Mock() caching_handler.request_kwargs = {"stream": True} - logging_obj._llm_caching_handler = caching_handler + logging_obj.llm_caching_handler = caching_handler iterator: Final = _make_iterator(sse_events=[], logging_obj=logging_obj) iterator.completed_response = ResponseCompletedEvent.model_construct( type="response.completed", @@ -1285,7 +1285,7 @@ def test_persist_completed_response_to_cache_skips_a_response_without_output(mon iterator._persist_completed_response_to_cache(is_async=False) cache.add_cache.assert_not_called() - caching_handler._should_store_result_in_cache.assert_not_called() + caching_handler.should_store_result_in_cache.assert_not_called() def test_persist_completed_response_to_cache_survives_an_unserializable_response(monkeypatch): @@ -1296,7 +1296,7 @@ def test_persist_completed_response_to_cache_survives_an_unserializable_response logging_obj: Final = _logging_obj_stub() caching_handler: Final = Mock() caching_handler.request_kwargs = {"stream": True} - logging_obj._llm_caching_handler = caching_handler + logging_obj.llm_caching_handler = caching_handler iterator: Final = _make_iterator(sse_events=[], logging_obj=logging_obj) iterator.completed_response = ResponseCompletedEvent.model_construct( type="response.completed", response=bad_response diff --git a/tests/unit/responses/test_streaming_iterator_error_events.py b/tests/unit/responses/test_streaming_iterator_error_events.py index e7cf09909fe..2286aadc894 100644 --- a/tests/unit/responses/test_streaming_iterator_error_events.py +++ b/tests/unit/responses/test_streaming_iterator_error_events.py @@ -452,7 +452,7 @@ def test_handle_logging_failed_response_records_usage_and_cost(): usage=usage, ) iterator.completed_response = chunk - iterator.logging_obj._response_cost_calculator.return_value = 0.0042 + iterator.logging_obj.response_cost_calculator.return_value = 0.0042 with ( patch.object(import_module("litellm.responses.streaming_iterator"), "run_async_function"), patch.object(import_module("litellm.responses.streaming_iterator"), "executor"), @@ -464,7 +464,7 @@ def test_handle_logging_failed_response_records_usage_and_cost(): assert combined_usage.completion_tokens == 5 assert combined_usage.total_tokens == 15 assert iterator.logging_obj.model_call_details["response_cost"] == 0.0042 - iterator.logging_obj._response_cost_calculator.assert_called_once_with(result=chunk.response) + iterator.logging_obj.response_cost_calculator.assert_called_once_with(result=chunk.response) def test_handle_logging_failed_response_without_usage_skips_recording(): @@ -478,7 +478,7 @@ def test_handle_logging_failed_response_without_usage_skips_recording(): ): iterator._handle_logging_failed_response() assert "combined_usage_object" not in iterator.logging_obj.model_call_details - iterator.logging_obj._response_cost_calculator.assert_not_called() + iterator.logging_obj.response_cost_calculator.assert_not_called() def test_sync_iterator_raises_mid_stream_fallback_on_rate_limit_error_event(): diff --git a/tests/unit/responses/test_text_format_conversion.py b/tests/unit/responses/test_text_format_conversion.py index c68ad16c4af..8cbffa898f3 100644 --- a/tests/unit/responses/test_text_format_conversion.py +++ b/tests/unit/responses/test_text_format_conversion.py @@ -153,7 +153,7 @@ class TestTextFormatConversion: import_module("litellm.responses.main").base_llm_http_handler, "response_api_handler", new=mock_handler, ): - litellm._turn_on_debug() + litellm.turn_on_debug() # Call aresponses with text_format parameter response = await litellm.aresponses( diff --git a/tests/unit/router_strategy/test_budget_limiter_hotpath.py b/tests/unit/router_strategy/test_budget_limiter_hotpath.py index a417b789397..4667ac67efa 100644 --- a/tests/unit/router_strategy/test_budget_limiter_hotpath.py +++ b/tests/unit/router_strategy/test_budget_limiter_hotpath.py @@ -293,7 +293,7 @@ def test_router_add_deployment_registers_deployment_budget(disable_budget_sync, ) ) - budget_limiter = router._get_router_deployment_budget_limiter() + budget_limiter = router.get_router_deployment_budget_limiter() assert budget_limiter is not None config = budget_limiter._get_budget_config_for_deployment("runtime-budget-deployment") assert config is not None diff --git a/tests/unit/router_strategy/test_router_routing_groups.py b/tests/unit/router_strategy/test_router_routing_groups.py index 534aea47885..e18c9948e39 100644 --- a/tests/unit/router_strategy/test_router_routing_groups.py +++ b/tests/unit/router_strategy/test_router_routing_groups.py @@ -595,9 +595,9 @@ def test_init_routing_groups_with_none_clears_state(): } ] ) - assert router._routing_groups + assert router.routing_groups router._init_routing_groups(None) - assert router._routing_groups == {} + assert router.routing_groups == {} assert router._model_to_group == {} assert router._group_selectors == {} @@ -814,7 +814,7 @@ def _single_latency_group(): def _assert_still_routes_with_original_group(router, selector): - assert list(router._routing_groups) == ["g1"] + assert list(router.routing_groups) == ["g1"] assert router._model_to_group == {"filtered-model": "g1"} assert router._group_selectors["g1"]["latency-based-routing"] is selector assert router._get_routing_context("filtered-model", None) == ("latency-based-routing", selector) @@ -855,7 +855,7 @@ def test_failed_routing_groups_update_does_not_poison_later_strategy_changes(mon router.update_settings(routing_strategy="least-busy") - assert list(router._routing_groups) == ["g1"] + assert list(router.routing_groups) == ["g1"] assert [g["group_name"] for g in router.get_settings()["routing_groups"]] == ["g1"] @@ -958,7 +958,7 @@ def test_replace_routing_groups_swaps_state_and_callbacks_in_one_step(monkeypatc ) ) - assert list(router._routing_groups) == ["g2", "g3"] + assert list(router.routing_groups) == ["g2", "g3"] assert router._model_to_group == {"other-model": "g2", "other-model-2": "g3"} assert router._group_selectors == {"g2": {"least-busy": new_selector}, "g3": {}} assert router._get_routing_context("other-model", None) == ("least-busy", new_selector) @@ -1556,7 +1556,7 @@ def _pin_choice_to(deployment_id): async def _call_and_get_cooldowns(router, model): - from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments + from litellm.router_utils.cooldown_handlers import async_get_cooldown_deployments with ( patch("litellm.router_strategy.simple_shuffle.random.choice", side_effect=_pin_choice_to("deploy-3")), @@ -1567,7 +1567,7 @@ async def _call_and_get_cooldowns(router, model): messages=[{"role": "user", "content": "hi"}], mock_response="litellm.RateLimitError", ) - return await _async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) + return await async_get_cooldown_deployments(litellm_router_instance=router, parent_otel_span=None) @pytest.mark.asyncio diff --git a/tests/unit/router_strategy/test_router_tag_routing.py b/tests/unit/router_strategy/test_router_tag_routing.py index d46b12a338f..57c0177f9b4 100644 --- a/tests/unit/router_strategy/test_router_tag_routing.py +++ b/tests/unit/router_strategy/test_router_tag_routing.py @@ -393,39 +393,39 @@ async def test_router_free_paid_tier_with_responses_api(): def test_get_tags_from_request_kwargs_none(): - from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs + from litellm.router_strategy.tag_based_routing import get_tags_from_request_kwargs # None request kwargs should safely return empty list - assert _get_tags_from_request_kwargs(None) == [] + assert get_tags_from_request_kwargs(None) == [] def test_get_tags_from_request_kwargs_various_inputs(): - from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs + from litellm.router_strategy.tag_based_routing import get_tags_from_request_kwargs # Direct "metadata" path - assert _get_tags_from_request_kwargs({"metadata": {"tags": ["free"]}}) == ["free"] - assert _get_tags_from_request_kwargs({"metadata": {"tags": []}}) == [] - assert _get_tags_from_request_kwargs({"metadata": {"tags": None}}) == [] - assert _get_tags_from_request_kwargs({"metadata": {}}) == [] - assert _get_tags_from_request_kwargs({"metadata": None}) == [] + assert get_tags_from_request_kwargs({"metadata": {"tags": ["free"]}}) == ["free"] + assert get_tags_from_request_kwargs({"metadata": {"tags": []}}) == [] + assert get_tags_from_request_kwargs({"metadata": {"tags": None}}) == [] + assert get_tags_from_request_kwargs({"metadata": {}}) == [] + assert get_tags_from_request_kwargs({"metadata": None}) == [] # Indirect via "litellm_params" - metadata inside - assert _get_tags_from_request_kwargs({"litellm_params": {"metadata": {"tags": ["paid"]}}}) == ["paid"] - assert _get_tags_from_request_kwargs({"litellm_params": {"metadata": None}}) == [] - assert _get_tags_from_request_kwargs({"litellm_params": {}}) == [] + assert get_tags_from_request_kwargs({"litellm_params": {"metadata": {"tags": ["paid"]}}}) == ["paid"] + assert get_tags_from_request_kwargs({"litellm_params": {"metadata": None}}) == [] + assert get_tags_from_request_kwargs({"litellm_params": {}}) == [] # Alternate metadata variable name: "litellm_metadata" - assert _get_tags_from_request_kwargs( + assert get_tags_from_request_kwargs( {"litellm_metadata": {"tags": ["alt"]}}, metadata_variable_name="litellm_metadata", ) == ["alt"] - assert _get_tags_from_request_kwargs( + assert get_tags_from_request_kwargs( {"litellm_params": {"litellm_metadata": {"tags": ["nested-alt"]}}}, metadata_variable_name="litellm_metadata", ) == ["nested-alt"] # No relevant keys present - assert _get_tags_from_request_kwargs({"foo": "bar"}) == [] + assert get_tags_from_request_kwargs({"foo": "bar"}) == [] @pytest.mark.parametrize( @@ -444,15 +444,15 @@ def test_get_tags_from_request_kwargs_reads_no_tags_from_a_non_dict_shape(reques """Metadata and `tags` are request-controlled, so a client can send either as a string, a list or null. Every shape that cannot hold string tags reads as untagged instead of raising, because callers run on the hot request path.""" - from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs + from litellm.router_strategy.tag_based_routing import get_tags_from_request_kwargs - assert _get_tags_from_request_kwargs(request_kwargs) == [] + assert get_tags_from_request_kwargs(request_kwargs) == [] def test_get_tags_from_request_kwargs_keeps_only_string_tags(): - from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs + from litellm.router_strategy.tag_based_routing import get_tags_from_request_kwargs - assert _get_tags_from_request_kwargs({"metadata": {"tags": ["free", 7, None, "paid"]}}) == ["free", "paid"] + assert get_tags_from_request_kwargs({"metadata": {"tags": ["free", 7, None, "paid"]}}) == ["free", "paid"] # --- _split_tags unit tests --- @@ -1168,7 +1168,7 @@ class _FakeRouterForChainOverride: def __init__(self, all_deployments): self._all_deployments = all_deployments - def _get_all_deployments(self, model_name): + def get_all_deployments(self, model_name): return self._all_deployments @@ -1215,7 +1215,7 @@ def test_chain_tag_filtering_override_falls_back_to_healthy_deployments_on_looku from litellm.router_strategy.tag_based_routing import _chain_tag_filtering_override class _BrokenRouter: - def _get_all_deployments(self, model_name): + def get_all_deployments(self, model_name): raise RuntimeError("model group not found") healthy_deployments = [{"model_info": {"enable_tag_filtering": False}}] @@ -2536,7 +2536,7 @@ async def test_plain_tag_exhaustion_with_universal_default_tag_raises_by_default router = _quality_high_cost_low_router() with patch( - "litellm.router._async_get_cooldown_deployments", + "litellm.router.async_get_cooldown_deployments", new=AsyncMock(return_value=["quality-high-1", "quality-high-2"]), ): with pytest.raises(Exception, match='Not allowed to access model due to tags configuration\\.') as exc_info: @@ -2561,7 +2561,7 @@ async def test_plain_tag_exhaustion_with_universal_default_tag_falls_open_when_a from unittest.mock import AsyncMock, patch with patch( - "litellm.router._async_get_cooldown_deployments", + "litellm.router.async_get_cooldown_deployments", new=AsyncMock(return_value=["quality-high-1", "quality-high-2"]), ): response = await router.acompletion( diff --git a/tests/unit/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py b/tests/unit/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py index 3a92aa221e5..5b204c4155b 100644 --- a/tests/unit/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py +++ b/tests/unit/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py @@ -94,7 +94,7 @@ class TestEncryptedItemIdCodec: original_item_id = "rs_abc123def456" encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_item_id) assert encoded.startswith("encitem_") - decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded) + decoded = ResponsesAPIRequestUtils.decode_encrypted_item_id(encoded) assert decoded is not None assert decoded["model_id"] == model_id assert decoded["item_id"] == original_item_id @@ -106,22 +106,22 @@ class TestEncryptedItemIdCodec: encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_item_id) # Strip any trailing '=' to simulate what happens in transit stripped = encoded.rstrip("=") - decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(stripped) + decoded = ResponsesAPIRequestUtils.decode_encrypted_item_id(stripped) assert decoded is not None assert decoded["model_id"] == model_id assert decoded["item_id"] == original_item_id def test_non_encoded_id_returns_none(self): - assert ResponsesAPIRequestUtils._decode_encrypted_item_id("rs_abc123") is None - assert ResponsesAPIRequestUtils._decode_encrypted_item_id("msg_abc") is None - assert ResponsesAPIRequestUtils._decode_encrypted_item_id("") is None + assert ResponsesAPIRequestUtils.decode_encrypted_item_id("rs_abc123") is None + assert ResponsesAPIRequestUtils.decode_encrypted_item_id("msg_abc") is None + assert ResponsesAPIRequestUtils.decode_encrypted_item_id("") is None def test_semicolon_in_item_id(self): """item_id values containing ';' must survive the roundtrip.""" model_id = "deployment-1" original_item_id = "rs_part1;part2;part3" encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_item_id) - decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded) + decoded = ResponsesAPIRequestUtils.decode_encrypted_item_id(encoded) assert decoded is not None assert decoded["item_id"] == original_item_id @@ -142,7 +142,7 @@ class TestUpdateEncryptedContentItemIds: # Reasoning item with encrypted_content gets encoded encoded_id = result["output"][1]["id"] assert encoded_id.startswith("encitem_") - decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded_id) + decoded = ResponsesAPIRequestUtils.decode_encrypted_item_id(encoded_id) assert decoded["model_id"] == model_id assert decoded["item_id"] == "rs_xyz" @@ -157,14 +157,14 @@ class TestEncryptedContentWrapping: """Test wrapping encrypted_content with model_id metadata.""" model_id = "deployment-1" original_content = "gAAAAABpnW_yEYmSNEyOG_original_encrypted_data" - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id) + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id(original_content, model_id) assert wrapped.startswith("litellm_enc:") assert wrapped != original_content ( unwrapped_model_id, unwrapped_content, - ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(wrapped) + ) = ResponsesAPIRequestUtils.unwrap_encrypted_content_with_model_id(wrapped) assert unwrapped_model_id == model_id assert unwrapped_content == original_content @@ -174,7 +174,7 @@ class TestEncryptedContentWrapping: ( model_id, content, - ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(plain_content) + ) = ResponsesAPIRequestUtils.unwrap_encrypted_content_with_model_id(plain_content) assert model_id is None assert content == plain_content @@ -199,7 +199,7 @@ class TestEncryptedContentWrapping: ( model_id_extracted, unwrapped, - ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(wrapped) + ) = ResponsesAPIRequestUtils.unwrap_encrypted_content_with_model_id(wrapped) assert model_id_extracted == model_id assert unwrapped == "gAAAAABpnW_yEYmSNEyOG_secret" @@ -214,7 +214,7 @@ class TestRestoreEncryptedContentItemIds: {"type": "message", "id": "msg_abc123", "role": "assistant"}, {"type": "reasoning", "id": encoded_id, "encrypted_content": "secret"}, ] - restored = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(request_input) + restored = ResponsesAPIRequestUtils.restore_encrypted_content_item_ids_in_input(request_input) assert restored[0]["id"] == "msg_abc123" assert restored[1]["id"] == original_id @@ -222,21 +222,21 @@ class TestRestoreEncryptedContentItemIds: """Test that wrapped encrypted_content is unwrapped before forwarding.""" model_id = "deployment-1" original_content = "gAAAAABpnW_yEYmSNEyOG_original" - wrapped_content = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id) + wrapped_content = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id(original_content, model_id) request_input = [ {"type": "reasoning", "encrypted_content": wrapped_content}, ] - restored = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(request_input) + restored = ResponsesAPIRequestUtils.restore_encrypted_content_item_ids_in_input(request_input) assert restored[0]["encrypted_content"] == original_content def test_no_op_for_plain_string_input(self): - result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input("Hello world") + result = ResponsesAPIRequestUtils.restore_encrypted_content_item_ids_in_input("Hello world") assert result == "Hello world" def test_no_op_for_unencoded_ids(self): request_input = [{"type": "message", "id": "msg_plain"}] - result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(request_input) + result = ResponsesAPIRequestUtils.restore_encrypted_content_item_ids_in_input(request_input) assert result[0]["id"] == "msg_plain" @@ -330,7 +330,7 @@ async def test_encrypted_content_affinity_tracks_and_routes(): ) # Verify the encoded ID decodes back to the correct deployment + original ID - decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded_item_id) + decoded = ResponsesAPIRequestUtils.decode_encrypted_item_id(encoded_item_id) assert decoded is not None assert decoded["model_id"] == first_model_id assert decoded["item_id"] == "rs_encrypted_item_456" @@ -610,7 +610,7 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id(): ( extracted_model_id, _, - ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(wrapped_content) + ) = ResponsesAPIRequestUtils.unwrap_encrypted_content_with_model_id(wrapped_content) assert extracted_model_id == first_model_id # Second request: use wrapped encrypted_content WITHOUT an ID (Codex behavior) @@ -638,7 +638,7 @@ def test_encrypted_content_wrapping_preserves_original_content(): model_id = "test-deployment-1" original_encrypted_content = "gAAAAABpnW_yEYmSNEyOG_streaming_test_content_with_special_chars==+/" - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_encrypted_content, model_id) + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id(original_encrypted_content, model_id) assert wrapped.startswith("litellm_enc:") assert wrapped != original_encrypted_content @@ -646,7 +646,7 @@ def test_encrypted_content_wrapping_preserves_original_content(): ( extracted_model_id, unwrapped_content, - ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(wrapped) + ) = ResponsesAPIRequestUtils.unwrap_encrypted_content_with_model_id(wrapped) assert extracted_model_id == model_id assert unwrapped_content == original_encrypted_content @@ -659,12 +659,12 @@ def test_encrypted_content_wrapping_with_multiple_semicolons(): model_id = "deployment-with-semicolons" original_content = "gAAAAAB;some;content;with;semicolons" - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id) + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id(original_content, model_id) ( extracted_model_id, unwrapped, - ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(wrapped) + ) = ResponsesAPIRequestUtils.unwrap_encrypted_content_with_model_id(wrapped) assert extracted_model_id == model_id assert unwrapped == original_content @@ -741,14 +741,14 @@ def test_encrypted_content_wrapping_empty_string(): model_id = "test-deployment" original_content = "" - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id) + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id(original_content, model_id) assert wrapped.startswith("litellm_enc:") ( extracted_model_id, unwrapped, - ) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(wrapped) + ) = ResponsesAPIRequestUtils.unwrap_encrypted_content_with_model_id(wrapped) assert extracted_model_id == model_id assert unwrapped == original_content @@ -1369,7 +1369,7 @@ async def test_affinity_strips_and_dispatches_when_origin_cooled_for_non_429(): { "id": encoded_id, "type": "reasoning", - "encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( + "encrypted_content": ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id( "gAAAAA-blob", "deployment-a-cooled" ), } @@ -1504,7 +1504,7 @@ async def test_affinity_serves_sibling_when_candidate_origin_has_no_boundary_pee originating, cooldown_entries=[], routed_group_model_ids=["region-a", "region-b", "region-c"] ) check = EncryptedContentAffinityCheck(router=mock_router) - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "region-a") + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA-blob", "region-a") siblings = [ { "model_info": {"id": "region-b"}, @@ -1566,7 +1566,7 @@ async def test_affinity_strips_and_dispatches_when_origin_is_unknown_or_removed( mock_router.get_deployment.return_value = None check = EncryptedContentAffinityCheck(router=mock_router) - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-removed") + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-removed") routed_pool = [ { "model_info": {"id": "deployment-b"}, @@ -1850,7 +1850,7 @@ async def test_encrypted_content_affinity_pins_anthropic_messages_replayed_throu def _bridge_replayed_anthropic_messages(minted_by: str) -> list: - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA_turn_one", minted_by) + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA_turn_one", minted_by) return [ {"role": "user", "content": "Solve the zebra puzzle"}, { @@ -1908,7 +1908,7 @@ async def test_encrypted_content_affinity_strips_bridge_reasoning_from_messages_ class TestStripEncryptedReasoningFromInput: def test_keeps_summary_and_drops_encrypted_content_and_id(self): - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a") + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a") encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a", "rs_1") request_input = [ {"role": "user", "content": "first turn"}, @@ -1935,7 +1935,7 @@ class TestStripEncryptedReasoningFromInput: ] def test_keeps_string_form_summary_when_stripping(self): - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a") + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a") request_input = [ {"type": "reasoning", "encrypted_content": wrapped, "summary": "plain string thought"}, { @@ -1962,7 +1962,7 @@ class TestStripEncryptedReasoningFromInput: assert request_input == before def test_strips_only_items_selected_by_predicate(self): - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a") + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a") request_input = [ {"type": "reasoning", "id": "keep", "encrypted_content": wrapped, "summary": "keep"}, {"type": "reasoning", "id": "strip", "encrypted_content": wrapped, "summary": "strip"}, @@ -2007,8 +2007,8 @@ async def test_real_router_selection_keeps_origin_reasoning_and_strips_foreign_o ) openai_item_id = ResponsesAPIRequestUtils._build_encrypted_item_id("dep-openai", "rs-openai") azure_item_id = ResponsesAPIRequestUtils._build_encrypted_item_id("dep-azure", "rs-azure") - openai_wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("blob-openai", "dep-openai") - azure_wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("blob-azure", "dep-azure") + openai_wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("blob-openai", "dep-openai") + azure_wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("blob-azure", "dep-azure") request_input = [ {"type": "message", "role": "user", "content": "first question"}, { @@ -2079,14 +2079,14 @@ async def test_affinity_keeps_mixed_origins_on_the_same_encryption_boundary(): } d2_item = { "type": "reasoning", - "encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("blob-d2", "d2"), + "encrypted_content": ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("blob-d2", "d2"), "summary": [{"type": "summary_text", "text": "second origin"}], } request_kwargs = { "input": [ { "type": "reasoning", - "encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("blob-d1", "d1"), + "encrypted_content": ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("blob-d1", "d1"), "summary": [{"type": "summary_text", "text": "first origin"}], }, d2_item.copy(), @@ -2124,14 +2124,14 @@ async def test_boundary_pin_strips_reasoning_from_a_different_origin(): "input": [ { "type": "reasoning", - "encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( + "encrypted_content": ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id( "blob-origin-a", "origin-a" ), "summary": [{"type": "summary_text", "text": "origin A summary"}], }, { "type": "reasoning", - "encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( + "encrypted_content": ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id( "blob-origin-b", "origin-b" ), "summary": [{"type": "summary_text", "text": "origin B summary"}], @@ -2151,7 +2151,7 @@ async def test_boundary_pin_strips_reasoning_from_a_different_origin(): assert request_kwargs["input"] == [ { "type": "reasoning", - "encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( + "encrypted_content": ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id( "blob-origin-a", "origin-a" ), "summary": [{"type": "summary_text", "text": "origin A summary"}], @@ -2199,7 +2199,7 @@ async def test_affinity_keeps_only_anthropic_reasoning_from_the_pinned_origin(): "type": "redacted_thinking", "data": ( "litellm_encrypted_reasoning:" - f"{ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id('gAAAAA_turn_one', 'origin-a')}" + f"{ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id('gAAAAA_turn_one', 'origin-a')}" ), }, { @@ -2207,7 +2207,7 @@ async def test_affinity_keeps_only_anthropic_reasoning_from_the_pinned_origin(): "thinking": "The bridge packed this one", "signature": ( "litellm_encrypted_reasoning:" - f"{ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id('gAAAAA_turn_one', 'origin-a')}" + f"{ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id('gAAAAA_turn_one', 'origin-a')}" ), }, {"type": "text", "text": "The zebra owner lives in the green house."}, @@ -2230,7 +2230,7 @@ async def test_affinity_strips_unknown_origins_but_leaves_unmarked_encrypted_con } openai_item = { "type": "reasoning", - "encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("blob-a", "origin-a"), + "encrypted_content": ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("blob-a", "origin-a"), "summary": [{"type": "summary_text", "text": "origin A"}], } request_kwargs = { @@ -2238,7 +2238,7 @@ async def test_affinity_strips_unknown_origins_but_leaves_unmarked_encrypted_con openai_item.copy(), { "type": "reasoning", - "encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( + "encrypted_content": ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id( "blob-removed", "origin-removed" ), "summary": [{"type": "summary_text", "text": "removed origin"}], @@ -2272,7 +2272,7 @@ async def test_affinity_strips_unknown_origins_but_leaves_unmarked_encrypted_con def _cross_group_request_kwargs(): - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a") + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a") return { "litellm_metadata": {}, "input": [ @@ -2405,7 +2405,7 @@ async def test_affinity_strips_when_group_is_spelled_differently_but_same_by_id( originating, cooldown_entries=[], routed_group_model_ids=["deployment-mini-a", "deployment-mini-b"] ) check = EncryptedContentAffinityCheck(router=mock_router) - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-mini-a") + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-mini-a") sibling_pool = [ { "model_info": {"id": "deployment-mini-b"}, @@ -2457,7 +2457,7 @@ async def test_affinity_strips_for_team_and_pattern_routes(): originating, cooldown_entries=[], routed_group_model_ids=["deployment-team-a", "deployment-team-b"] ) check = EncryptedContentAffinityCheck(router=mock_router) - wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-team-a") + wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-team-a") sibling_pool = [ { "model_info": {"id": "deployment-team-b"}, diff --git a/tests/unit/router_utils/pre_call_checks/test_responses_api_deployment_check.py b/tests/unit/router_utils/pre_call_checks/test_responses_api_deployment_check.py index 78cafbec70a..64929bebaf5 100644 --- a/tests/unit/router_utils/pre_call_checks/test_responses_api_deployment_check.py +++ b/tests/unit/router_utils/pre_call_checks/test_responses_api_deployment_check.py @@ -102,7 +102,7 @@ async def test_async_responses_api_routing_with_previous_response_id(): mock_post.return_value = MockResponse(mock_response_data, 200) # Make the initial request - # litellm._turn_on_debug() + # litellm.turn_on_debug() response = await router.aresponses( model=MODEL, input="Hello, how are you?", diff --git a/tests/unit/router_utils/test_cooldown_cache.py b/tests/unit/router_utils/test_cooldown_cache.py index 7a8b9c5d15a..98f2b254a3f 100644 --- a/tests/unit/router_utils/test_cooldown_cache.py +++ b/tests/unit/router_utils/test_cooldown_cache.py @@ -240,7 +240,7 @@ class TestCooldownCacheExceptionMasking: # Test masking behavior with these settings long_string = "A" * 100 # 100 character string - masked = cache.exception_masker._mask_value(long_string) + masked = cache.exception_masker.mask_value(long_string) # Should show first 50 characters, then all asterisks expected = "A" * 50 + "*" * 50 diff --git a/tests/unit/router_utils/test_fallback_event_handlers.py b/tests/unit/router_utils/test_fallback_event_handlers.py index ec05a2c2840..f83dfa131ed 100644 --- a/tests/unit/router_utils/test_fallback_event_handlers.py +++ b/tests/unit/router_utils/test_fallback_event_handlers.py @@ -800,7 +800,7 @@ class TestTriggerCooldownForFailedDeployment: exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") exc.failed_deployment_id = "fallback-deployment" - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) mock_set_cooldown.assert_called_once() @@ -825,7 +825,7 @@ class TestTriggerCooldownForFailedDeployment: } } - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs=kwargs, exception=exc) mock_set_cooldown.assert_not_called() @@ -842,7 +842,7 @@ class TestTriggerCooldownForFailedDeployment: exc.failed_deployment_id = "fallback-deployment" with ( - patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown, + patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown, patch( "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" ) as mock_increment, @@ -857,7 +857,7 @@ class TestTriggerCooldownForFailedDeployment: def test_no_op_when_deployment_id_missing(self): mock_router = MagicMock() - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment( litellm_router=mock_router, kwargs={}, exception=RuntimeError("no metadata") ) @@ -873,7 +873,7 @@ class TestTriggerCooldownForFailedDeployment: exc.failed_deployment_id = "fallback-deployment" mark_advisor_orchestration_failure(exc) - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) mock_set_cooldown.assert_not_called() @@ -886,7 +886,7 @@ class TestTriggerCooldownForFailedDeployment: exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") exc.failed_deployment_id = "fallback-deployment" - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) call_kwargs = mock_set_cooldown.call_args[1] @@ -904,7 +904,7 @@ class TestTriggerCooldownForFailedDeployment: exc.failed_deployment_id = "fallback-deployment" exc.litellm_response_headers = httpx.Headers({"retry-after": "45"}) - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) call_kwargs = mock_set_cooldown.call_args[1] @@ -919,7 +919,7 @@ class TestTriggerCooldownForFailedDeployment: exc.failed_deployment_id = "fallback-deployment" with patch( - "litellm.router_utils.fallback_event_handlers._set_cooldown_deployments", + "litellm.router_utils.fallback_event_handlers.set_cooldown_deployments", side_effect=RuntimeError("cooldown error"), ): _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) @@ -937,7 +937,7 @@ class TestTriggerCooldownForFailedDeployment: exc.failed_deployment_id = "fallback-deployment" with ( - patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown, + patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown, patch( "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" ) as mock_increment, @@ -962,7 +962,7 @@ class TestTriggerCooldownForFailedDeployment: exc = litellm.NotFoundError("not found", "openai", "gpt-4") exc.failed_deployment_id = "fallback-deployment" - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) mock_set_cooldown.assert_called_once() @@ -984,7 +984,7 @@ class TestTriggerCooldownForFailedDeployment: exc.failed_deployment_id = "fallback-deployment" with ( - patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown, + patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown, patch( "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" ) as mock_increment, @@ -1059,7 +1059,7 @@ class TestTriggerCooldownForFailedDeployment: exc = litellm.Timeout(message="timeout", model="gpt-4", llm_provider="openai") exc.failed_deployment_id = "fallback-deployment" - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + with patch("litellm.router_utils.fallback_event_handlers.set_cooldown_deployments") as mock_set_cooldown: _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) mock_set_cooldown.assert_called_once() diff --git a/tests/unit/router_utils/test_health_check_allowed_fails_integration.py b/tests/unit/router_utils/test_health_check_allowed_fails_integration.py index 9021d842daa..d200b304c4e 100644 --- a/tests/unit/router_utils/test_health_check_allowed_fails_integration.py +++ b/tests/unit/router_utils/test_health_check_allowed_fails_integration.py @@ -54,9 +54,7 @@ class TestHealthCheckEndpointExceptionPropagation: from litellm.proxy.health_check import _perform_health_check - auth_error = litellm.AuthenticationError( - message="Invalid key", llm_provider="openai", model="gpt-4" - ) + auth_error = litellm.AuthenticationError(message="Invalid key", llm_provider="openai", model="gpt-4") model_list = [ { "model_name": "gpt-4", @@ -67,9 +65,7 @@ class TestHealthCheckEndpointExceptionPropagation: with patch( "litellm.proxy.health_check.litellm.ahealth_check", - new=AsyncMock( - return_value={"error": "auth failed", "exception": auth_error} - ), + new=AsyncMock(return_value={"error": "auth failed", "exception": auth_error}), ): healthy, unhealthy, exc_map = await _perform_health_check(model_list) @@ -85,9 +81,7 @@ class TestHealthCheckEndpointExceptionPropagation: from litellm.proxy.health_check import _perform_health_check - raw_exc = litellm.RateLimitError( - message="Rate limited", llm_provider="openai", model="gpt-4" - ) + raw_exc = litellm.RateLimitError(message="Rate limited", llm_provider="openai", model="gpt-4") model_list = [ { "model_name": "gpt-4", @@ -125,18 +119,14 @@ class TestGetAllowedFailsFromPolicyWithHealthCheckExceptions: (litellm.BadRequestError, "BadRequestErrorAllowedFails", 7), ], ) - def test_policy_resolves_for_health_check_exception_types( - self, exception_type, policy_field, threshold - ): + def test_policy_resolves_for_health_check_exception_types(self, exception_type, policy_field, threshold): """Each exception type from a health check should resolve to its policy threshold.""" policy = AllowedFailsPolicy(**{policy_field: threshold}) router = Router( model_list=[_make_model("d1")], allowed_fails_policy=policy, ) - exception = exception_type( - message="health check failed", llm_provider="openai", model="gpt-4" - ) + exception = exception_type(message="health check failed", llm_provider="openai", model="gpt-4") result = router.get_allowed_fails_from_policy(exception=exception) assert result == threshold @@ -153,7 +143,7 @@ class TestGetAllowedFailsFromPolicyWithHealthCheckExceptions: class TestHealthCheckCooldownIntegration: - """Test that health check failures trigger cooldown via _set_cooldown_deployments.""" + """Test that health check failures trigger cooldown via set_cooldown_deployments.""" def test_health_check_failure_increments_failed_calls(self): """Health check failure should increment the failed_calls counter.""" @@ -166,9 +156,7 @@ class TestHealthCheckCooldownIntegration: allowed_fails_policy=AllowedFailsPolicy(TimeoutErrorAllowedFails=3), ) - timeout_exc = litellm.Timeout( - message="Health check timeout", model="gpt-4", llm_provider="openai" - ) + timeout_exc = litellm.Timeout(message="Health check timeout", model="gpt-4", llm_provider="openai") # First call: should not cooldown (1 <= 3) result = should_cooldown_based_on_allowed_fails_policy( @@ -193,9 +181,7 @@ class TestHealthCheckCooldownIntegration: allowed_fails_policy=AllowedFailsPolicy(AuthenticationErrorAllowedFails=2), ) - auth_exc = litellm.AuthenticationError( - message="Invalid key", model="gpt-4", llm_provider="openai" - ) + auth_exc = litellm.AuthenticationError(message="Invalid key", model="gpt-4", llm_provider="openai") # Fails 1 and 2: should not cooldown for _ in range(2): @@ -249,7 +235,7 @@ class TestHealthCheckCooldownIntegration: def test_healthy_endpoints_do_not_trigger_cooldown(self): """Healthy endpoints should not increment any failure counters.""" - from litellm.router_utils.cooldown_handlers import _set_cooldown_deployments + from litellm.router_utils.cooldown_handlers import set_cooldown_deployments router = Router( model_list=[_make_model("deploy-1")], @@ -268,7 +254,7 @@ class TestHealthCheckCooldownIntegration: def test_disable_cooldowns_prevents_health_check_cooldown(self): """When disable_cooldowns=True, health check failures should not trigger cooldown.""" - from litellm.router_utils.cooldown_handlers import _set_cooldown_deployments + from litellm.router_utils.cooldown_handlers import set_cooldown_deployments router = Router( model_list=[_make_model("deploy-1"), _make_model("deploy-2", "gpt-5")], @@ -277,11 +263,9 @@ class TestHealthCheckCooldownIntegration: disable_cooldowns=True, ) - timeout_exc = litellm.Timeout( - message="Health check timeout", model="gpt-4", llm_provider="openai" - ) + timeout_exc = litellm.Timeout(message="Health check timeout", model="gpt-4", llm_provider="openai") - result = _set_cooldown_deployments( + result = set_cooldown_deployments( litellm_router_instance=router, original_exception=timeout_exc, exception_status=500, @@ -295,7 +279,7 @@ class TestWriteHealthStateIntegration: """Test _write_health_state_to_router_cache integrates with cooldown pipeline.""" def test_unhealthy_endpoint_triggers_set_cooldown(self): - """_write_health_state_to_router_cache should call _set_cooldown_deployments for unhealthy endpoints.""" + """_write_health_state_to_router_cache should call set_cooldown_deployments for unhealthy endpoints.""" import litellm.proxy.proxy_server as proxy_module from litellm.proxy.proxy_server import _write_health_state_to_router_cache @@ -305,9 +289,7 @@ class TestWriteHealthStateIntegration: enable_health_check_routing=True, ) - timeout_exc = litellm.Timeout( - message="Health check timeout", model="", llm_provider="" - ) + timeout_exc = litellm.Timeout(message="Health check timeout", model="", llm_provider="") unhealthy_endpoints = [ {"model_id": "deploy-1", "error": "timeout"}, @@ -317,9 +299,7 @@ class TestWriteHealthStateIntegration: ] with patch.object(proxy_module, "llm_router", router): - with patch( - "litellm.router_utils.cooldown_handlers._set_cooldown_deployments" - ) as mock_cooldown: + with patch("litellm.router_utils.cooldown_handlers.set_cooldown_deployments") as mock_cooldown: _write_health_state_to_router_cache( healthy_endpoints=healthy_endpoints, unhealthy_endpoints=unhealthy_endpoints, @@ -349,9 +329,7 @@ class TestWriteHealthStateIntegration: ] with patch.object(proxy_module, "llm_router", router): - with patch( - "litellm.router_utils.cooldown_handlers._set_cooldown_deployments" - ) as mock_cooldown: + with patch("litellm.router_utils.cooldown_handlers.set_cooldown_deployments") as mock_cooldown: _write_health_state_to_router_cache( healthy_endpoints=[], unhealthy_endpoints=unhealthy_endpoints, @@ -370,9 +348,7 @@ class TestWriteHealthStateIntegration: enable_health_check_routing=True, ) - rate_exc = litellm.RateLimitError( - message="Rate limited", model="gpt-4", llm_provider="openai" - ) + rate_exc = litellm.RateLimitError(message="Rate limited", model="gpt-4", llm_provider="openai") unhealthy_endpoints = [ {"model_id": "deploy-1", "error": "rate limited"}, @@ -382,9 +358,7 @@ class TestWriteHealthStateIntegration: with patch( "litellm.router_utils.router_callbacks.track_deployment_metrics.increment_deployment_failures_for_current_minute" ) as mock_increment: - with patch( - "litellm.router_utils.cooldown_handlers._set_cooldown_deployments" - ): + with patch("litellm.router_utils.cooldown_handlers.set_cooldown_deployments"): _write_health_state_to_router_cache( healthy_endpoints=[], unhealthy_endpoints=unhealthy_endpoints, @@ -433,9 +407,7 @@ class TestHealthCheckFilterBypassWithPolicy: # Filter should pass all through because policy is set result = router._filter_health_check_unhealthy_deployments(deployments) - assert ( - len(result) == 2 - ), "Binary filter should be bypassed when allowed_fails_policy is set" + assert len(result) == 2, "Binary filter should be bypassed when allowed_fails_policy is set" def test_filter_active_when_no_policy(self): """Binary health check filter still works when no allowed_fails_policy is configured.""" @@ -497,9 +469,7 @@ class TestHealthCheckFilterBypassWithPolicy: deployments = [_make_model("deploy-1"), _make_model("deploy-2", "gpt-5")] - result = await router._async_filter_health_check_unhealthy_deployments( - deployments - ) + result = await router._async_filter_health_check_unhealthy_deployments(deployments) assert len(result) == 2 def _make_scoped_router_with_unhealthy(self, policy) -> Router: @@ -535,9 +505,7 @@ class TestHealthCheckFilterBypassWithPolicy: def test_filter_with_policy_still_applies_to_listed_groups(self): """A model-group allowlist keeps the filter active for listed groups even with a policy set.""" - router = self._make_scoped_router_with_unhealthy( - AllowedFailsPolicy(AuthenticationErrorAllowedFails=3) - ) + router = self._make_scoped_router_with_unhealthy(AllowedFailsPolicy(AuthenticationErrorAllowedFails=3)) deployments = [ _make_model("bad-listed"), _make_model("ok-listed"), @@ -550,18 +518,14 @@ class TestHealthCheckFilterBypassWithPolicy: @pytest.mark.asyncio async def test_async_filter_with_policy_still_applies_to_listed_groups(self): """Async version: listed groups stay filtered with a policy set, unlisted stay untouched.""" - router = self._make_scoped_router_with_unhealthy( - AllowedFailsPolicy(TimeoutErrorAllowedFails=2) - ) + router = self._make_scoped_router_with_unhealthy(AllowedFailsPolicy(TimeoutErrorAllowedFails=2)) deployments = [ _make_model("bad-listed"), _make_model("ok-listed"), _make_model("bad-unlisted", "gpt-5"), ] - result = await router._async_filter_health_check_unhealthy_deployments( - deployments - ) + result = await router._async_filter_health_check_unhealthy_deployments(deployments) assert [d["model_info"]["id"] for d in result] == ["ok-listed", "bad-unlisted"] @@ -605,7 +569,7 @@ class TestAllDeploymentsInCooldownSafetyNet: # Simulate all deployments in cooldown with patch( - "litellm.router._async_get_cooldown_deployments", + "litellm.router.async_get_cooldown_deployments", new=AsyncMock(return_value=["deploy-1", "deploy-2"]), ): # The safety net in async_get_available_deployment should restore @@ -620,9 +584,7 @@ class TestAllDeploymentsInCooldownSafetyNet: if not filtered and router.enable_health_check_routing: filtered = _pre - assert ( - len(filtered) == 2 - ), "Safety net should return all deployments when all are in cooldown" + assert len(filtered) == 2, "Safety net should return all deployments when all are in cooldown" class TestHealthCheckIgnoreTransientErrors: @@ -644,9 +606,7 @@ class TestHealthCheckIgnoreTransientErrors: health_check_ignore_transient_errors=True, ) - rate_exc = litellm.RateLimitError( - message="Rate limited", model="gpt-4", llm_provider="openai" - ) + rate_exc = litellm.RateLimitError(message="Rate limited", model="gpt-4", llm_provider="openai") assert getattr(rate_exc, "status_code", None) == 429 unhealthy_endpoints = [ @@ -654,9 +614,7 @@ class TestHealthCheckIgnoreTransientErrors: ] with patch.object(proxy_module, "llm_router", router): - with patch( - "litellm.router_utils.cooldown_handlers._set_cooldown_deployments" - ) as mock_cooldown: + with patch("litellm.router_utils.cooldown_handlers.set_cooldown_deployments") as mock_cooldown: with patch( "litellm.router_utils.router_callbacks.track_deployment_metrics.increment_deployment_failures_for_current_minute" ) as mock_increment: @@ -680,18 +638,14 @@ class TestHealthCheckIgnoreTransientErrors: health_check_ignore_transient_errors=True, ) - timeout_exc = litellm.Timeout( - message="Health check timeout exceeded", model="", llm_provider="" - ) + timeout_exc = litellm.Timeout(message="Health check timeout exceeded", model="", llm_provider="") unhealthy_endpoints = [ {"model_id": "deploy-1", "error": "timeout"}, ] with patch.object(proxy_module, "llm_router", router): - with patch( - "litellm.router_utils.cooldown_handlers._set_cooldown_deployments" - ) as mock_cooldown: + with patch("litellm.router_utils.cooldown_handlers.set_cooldown_deployments") as mock_cooldown: _write_health_state_to_router_cache( healthy_endpoints=[], unhealthy_endpoints=unhealthy_endpoints, @@ -711,18 +665,14 @@ class TestHealthCheckIgnoreTransientErrors: health_check_ignore_transient_errors=True, ) - auth_exc = litellm.AuthenticationError( - message="Invalid key", model="gpt-4", llm_provider="openai" - ) + auth_exc = litellm.AuthenticationError(message="Invalid key", model="gpt-4", llm_provider="openai") unhealthy_endpoints = [ {"model_id": "deploy-1", "error": "auth failed"}, ] with patch.object(proxy_module, "llm_router", router): - with patch( - "litellm.router_utils.cooldown_handlers._set_cooldown_deployments" - ) as mock_cooldown: + with patch("litellm.router_utils.cooldown_handlers.set_cooldown_deployments") as mock_cooldown: _write_health_state_to_router_cache( healthy_endpoints=[], unhealthy_endpoints=unhealthy_endpoints, @@ -742,9 +692,7 @@ class TestHealthCheckIgnoreTransientErrors: health_check_ignore_transient_errors=True, ) - rate_exc = litellm.RateLimitError( - message="Rate limited", model="gpt-4", llm_provider="openai" - ) + rate_exc = litellm.RateLimitError(message="Rate limited", model="gpt-4", llm_provider="openai") unhealthy_endpoints = [ {"model_id": "deploy-1", "error": "rate limited"}, @@ -774,18 +722,14 @@ class TestHealthCheckIgnoreTransientErrors: health_check_ignore_transient_errors=False, ) - rate_exc = litellm.RateLimitError( - message="Rate limited", model="gpt-4", llm_provider="openai" - ) + rate_exc = litellm.RateLimitError(message="Rate limited", model="gpt-4", llm_provider="openai") unhealthy_endpoints = [ {"model_id": "deploy-1", "error": "rate limited"}, ] with patch.object(proxy_module, "llm_router", router): - with patch( - "litellm.router_utils.cooldown_handlers._set_cooldown_deployments" - ) as mock_cooldown: + with patch("litellm.router_utils.cooldown_handlers.set_cooldown_deployments") as mock_cooldown: _write_health_state_to_router_cache( healthy_endpoints=[], unhealthy_endpoints=unhealthy_endpoints, diff --git a/tests/unit/router_utils/test_reasoning_effort_capability.py b/tests/unit/router_utils/test_reasoning_effort_capability.py index 2f23320a8ad..a08c99e269c 100644 --- a/tests/unit/router_utils/test_reasoning_effort_capability.py +++ b/tests/unit/router_utils/test_reasoning_effort_capability.py @@ -111,9 +111,9 @@ class TestBareModelNameFallback: """azure/gpt-5-mini carries no effort flag while gpt-5-mini carries three, and the request path resolves capability flags through that same twin (#20885). Reading only the prefixed entry would answer unknown for a model the map fully describes.""" - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper - model_info = dict(_get_model_info_helper(model="gpt-5-mini", custom_llm_provider="azure")) + model_info = dict(get_model_info_helper(model="gpt-5-mini", custom_llm_provider="azure")) assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ( "minimal", @@ -194,9 +194,9 @@ class TestNoneLevelPolarity: """AzureOpenAIGPT5Config raises UnsupportedParamsError on reasoning_effort='none' for models it does not flag, so advertising the level there would offer routing a 400.""" from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper - model_info = dict(_get_model_info_helper(model=model_key.split("/", 1)[1], custom_llm_provider="azure")) + model_info = dict(get_model_info_helper(model=model_key.split("/", 1)[1], custom_llm_provider="azure")) resolved = resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) assert resolved is not None @@ -353,9 +353,9 @@ class TestKimiK3AdvertisesItsDocumentedLevels: def test_the_declaration_survives_model_info_hydration(self, local_model_cost_map, model, provider): """The hydration line is the load-bearing seam: without it the key the map carries never reaches the resolver and reads as absent everywhere downstream.""" - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper - model_info = dict(_get_model_info_helper(model=model, custom_llm_provider=provider)) + model_info = dict(get_model_info_helper(model=model, custom_llm_provider=provider)) assert model_info["reasoning_effort_levels"] == ["low", "high", "max"] assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ("low", "high", "max") @@ -378,9 +378,9 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels: def test_the_entry_advertises_low_through_max_without_none(self, local_model_cost_map): """OpenAI documents low, medium, high, xhigh and max for gpt-6-astra. Unlike gpt-5.6-sol it does not take none, so a group must not offer none and must offer max.""" - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper - model_info = dict(_get_model_info_helper(model="gpt-6-astra", custom_llm_provider="openai")) + model_info = dict(get_model_info_helper(model="gpt-6-astra", custom_llm_provider="openai")) assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ( "low", @@ -405,9 +405,9 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels: on both Azure routes: none returns 200 with zero reasoning tokens and unlocks temperature, which OpenAI's API rejects, while max returns 400 unsupported_value naming none through xhigh as the levels it does take.""" - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper - model_info = dict(_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)) + model_info = dict(get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)) assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ( "none", @@ -423,9 +423,9 @@ class TestGpt6SolAndLunaAdvertiseNoneThroughMax: def test_the_entry_advertises_none_through_max(self, local_model_cost_map, model): """OpenAI documents none, low, medium (default), high, xhigh and max for both. Unlike gpt-6-astra they take none.""" - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper - model_info = dict(_get_model_info_helper(model=model, custom_llm_provider="openai")) + model_info = dict(get_model_info_helper(model=model, custom_llm_provider="openai")) assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == ( "none", @@ -477,10 +477,10 @@ class TestAzureGpt6SolAndLunaAdvertiseTheOpenAiLevels: ): """The Foundry deployments of sol and luna take the same effort set OpenAI documents for the direct API, so the resolved levels must match the OpenAI-direct entry.""" - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper - azure_info = dict(_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)) - openai_info = dict(_get_model_info_helper(model=model.rsplit("/", 1)[1], custom_llm_provider="openai")) + azure_info = dict(get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)) + openai_info = dict(get_model_info_helper(model=model.rsplit("/", 1)[1], custom_llm_provider="openai")) assert resolve_supported_reasoning_efforts( azure_info, deployment_is_mapped=True diff --git a/tests/unit/router_utils/test_router_utils_common_utils.py b/tests/unit/router_utils/test_router_utils_common_utils.py index ac18b4889dd..e0829e42557 100644 --- a/tests/unit/router_utils/test_router_utils_common_utils.py +++ b/tests/unit/router_utils/test_router_utils_common_utils.py @@ -402,7 +402,7 @@ def test_filter_deployments_by_model_access_groups_access_group_only_key(): filtered = router._filter_deployments_by_model_access_groups( model="gpt-5", - healthy_deployments=router._get_all_deployments(model_name="gpt-5"), + healthy_deployments=router.get_all_deployments(model_name="gpt-5"), request_kwargs={ "metadata": { "user_api_key_team_id": "team-2", @@ -561,9 +561,9 @@ class TestResolveModelGroupAlias: model_group_alias={"group-a": "group-b", "group-item": {"model": "group-b", "hidden": True}}, ) - assert router._get_model_from_alias("group-a") == "group-b" - assert router._get_model_from_alias("group-item") == "group-b" - assert router._get_model_from_alias("group-b") is None + assert router.get_model_from_alias("group-a") == "group-b" + assert router.get_model_from_alias("group-item") == "group-b" + assert router.get_model_from_alias("group-b") is None class TestTruncateFallbackErrorDetail: diff --git a/tests/unit/router_utils/test_routing_read_batch.py b/tests/unit/router_utils/test_routing_read_batch.py index a74e3e24f99..beafe1e8603 100644 --- a/tests/unit/router_utils/test_routing_read_batch.py +++ b/tests/unit/router_utils/test_routing_read_batch.py @@ -44,7 +44,7 @@ def _router(redis: MagicMock, routing_strategy: str) -> Router: model_list=[_deployment("dep-a"), _deployment("dep-b")], routing_strategy=routing_strategy, ) - router._update_redis_cache(cache=redis) + router.update_redis_cache(cache=redis) return router diff --git a/tests/unit/rust_bridge/test_logger.py b/tests/unit/rust_bridge/test_logger.py index aff385b9344..5a7b6c270a1 100644 --- a/tests/unit/rust_bridge/test_logger.py +++ b/tests/unit/rust_bridge/test_logger.py @@ -15,8 +15,8 @@ from litellm._logging import ( from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH from litellm.litellm_core_utils.secret_redaction import ( _python_redact_internal_details, - _python_redact_string, - _python_redact_structured_value, + python_redact_string, + python_redact_structured_value, ) from litellm.rust_bridge import diagnostics, logger @@ -84,8 +84,8 @@ def test_native_credential_patterns_match_python(text: str) -> None: from litellm.rust_bridge._native import NativeDiagnosticProcessor processor: Final = NativeDiagnosticProcessor(MINIMUM_CUSTOM_KEY_LENGTH) - assert processor.redact_text(text) == _python_redact_string(text) - assert processor.redact_structured_text("api_key", "secret123") == _python_redact_structured_value( + assert processor.redact_text(text) == python_redact_string(text) + assert processor.redact_structured_text("api_key", "secret123") == python_redact_structured_value( "api_key", "secret123" ) diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 7d9b7b73aca..9e348c0579a 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -729,7 +729,7 @@ def test_custom_pricing_cost_calc_uses_router_model_id_from_litellm_metadata(): This tests the full chain that was broken for /messages and /responses endpoints. Regression test for #23185. """ - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model custom_model_id = "claude-sonnet-4-custom-pricing-test" @@ -756,7 +756,7 @@ def test_custom_pricing_cost_calc_uses_router_model_id_from_litellm_metadata(): # _select_model_name_for_cost_calc appends provider prefix to the # selected router_model_id, so the result is "anthropic/" - selected_model = _select_model_name_for_cost_calc( + selected_model = select_model_name_for_cost_calc( model="anthropic/claude-sonnet-4-20250514", completion_response=None, custom_pricing=custom_pricing, @@ -767,7 +767,7 @@ def test_custom_pricing_cost_calc_uses_router_model_id_from_litellm_metadata(): assert custom_model_id in selected_model # Without custom_pricing, the router_model_id is NOT selected - selected_model_no_custom = _select_model_name_for_cost_calc( + selected_model_no_custom = select_model_name_for_cost_calc( model="anthropic/claude-sonnet-4-20250514", completion_response=None, custom_pricing=False, @@ -787,7 +787,7 @@ def test_per_request_custom_pricing_with_router(): returned 0.0 for per-request custom pricing via Router. """ from litellm import Router - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc router = Router( model_list=[ @@ -824,7 +824,7 @@ def test_per_request_custom_pricing_with_router(): # _select_model_name_for_cost_calc should pick the model name (which has pricing), # NOT the router_model_id (which has no pricing) - selected = _select_model_name_for_cost_calc( + selected = select_model_name_for_cost_calc( model="openai/gpt-3.5-turbo", completion_response=None, custom_pricing=True, @@ -844,7 +844,7 @@ def test_tiered_pricing_only_deployment_selects_router_model_id(): (e.g. dashscope/qwen3.7-plus) being billed as free. """ from litellm import Router - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc router = Router( model_list=[ @@ -874,7 +874,7 @@ def test_tiered_pricing_only_deployment_selects_router_model_id(): # The stripped shared alias must not carry tiered pricing. assert litellm.model_cost["dashscope/qwen-tier-only-test"].get("tiered_pricing") is None - selected = _select_model_name_for_cost_calc( + selected = select_model_name_for_cost_calc( model="dashscope/qwen-tier-only-test", completion_response=None, custom_pricing=True, @@ -2251,7 +2251,7 @@ def test_completion_cost_anthropic_auto_tier_uses_served_priority_rate(_local_mo Priority is opted into with ``service_tier="auto"``; Anthropic then serves "priority" and reports it on the response usage. The proxy forwards the - request-level "auto" into ``completion_cost`` (via ``_response_cost_calculator``), + request-level "auto" into ``completion_cost`` (via ``response_cost_calculator``), and that preference must not shadow the served tier, otherwise priority requests are silently billed at the standard rate. """ @@ -2343,7 +2343,7 @@ def test_completion_cost_non_string_service_tier_defers_to_served_tier(_local_mo ``allowed_openai_params``/``drop_params``) must not crash cost tracking. Before the fix, ``completion_cost`` called ``service_tier.lower()`` on the - request-level value, so a dict raised ``AttributeError``. ``_response_cost_calculator`` + request-level value, so a dict raised ``AttributeError``. ``response_cost_calculator`` swallowed it and reported ``response_cost=None``, silently dropping the cost. The non-string preference must be ignored so pricing defers to the tier the provider actually served on the response usage. @@ -3827,7 +3827,7 @@ def test_select_model_name_strips_unregistered_alias_prefix(_local_model_cost_ma silently pricing every streamed request at $0. """ - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc response = litellm.ModelResponse( id="x", @@ -3842,7 +3842,7 @@ def test_select_model_name_strips_unregistered_alias_prefix(_local_model_cost_ma ) response._hidden_params = {} - selected = _select_model_name_for_cost_calc( + selected = select_model_name_for_cost_calc( model=None, completion_response=response, custom_llm_provider="vertex_ai", @@ -3855,7 +3855,7 @@ def test_select_model_name_strips_duplicated_region_segment(_local_model_cost_ma """A "region/model" alias whose leading segment repeats the request's region must resolve to the region-priced cost key instead of keeping the region segment twice.""" - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc response = litellm.ModelResponse( id="x", @@ -3870,7 +3870,7 @@ def test_select_model_name_strips_duplicated_region_segment(_local_model_cost_ma ) response._hidden_params = {"region_name": "us-east-1"} - selected = _select_model_name_for_cost_calc( + selected = select_model_name_for_cost_calc( model=None, completion_response=response, custom_llm_provider="bedrock", @@ -3899,9 +3899,9 @@ def test_select_model_name_applies_region_to_private_provider_response_model(_lo """A Bedrock stream carries its requested model as the private provider model and must keep the request's region in the cost key, exactly as the same request does without streaming.""" - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc - selected = _select_model_name_for_cost_calc( + selected = select_model_name_for_cost_calc( model=None, completion_response=_bedrock_response_with_private_model("anthropic.claude-v2:1", "us-east-1"), custom_llm_provider="bedrock", @@ -4126,9 +4126,9 @@ def test_select_model_name_keeps_base_model_free_of_region(_local_model_cost_map """An explicit base_model keeps pricing on that model's own key even when the request carries a region with different regional rates, so the private provider model never widens region pricing.""" - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc - selected = _select_model_name_for_cost_calc( + selected = select_model_name_for_cost_calc( model="my-bedrock-deployment", completion_response=_bedrock_response_with_private_model("moonshotai.kimi-k2.5", "ap-northeast-1"), base_model="moonshotai.kimi-k2.5", @@ -4164,7 +4164,7 @@ def test_completion_cost_base_model_ignores_regional_row(_local_model_cost_map): def test_select_model_name_unresolvable_alias_unchanged(_local_model_cost_map): """An alias that resolves to no known cost key keeps the legacy double-prefixed name.""" - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc response = litellm.ModelResponse( id="x", @@ -4179,7 +4179,7 @@ def test_select_model_name_unresolvable_alias_unchanged(_local_model_cost_map): ) response._hidden_params = {} - selected = _select_model_name_for_cost_calc( + selected = select_model_name_for_cost_calc( model=None, completion_response=response, custom_llm_provider="vertex_ai", @@ -4192,7 +4192,7 @@ def test_completion_cost_keeps_custom_priced_slash_router_id(_local_model_cost_m """A custom-priced router id containing "/" keeps its custom pricing instead of being rewritten to the built-in key its suffix happens to match.""" - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc litellm.register_model( model_cost={ @@ -4204,7 +4204,7 @@ def test_completion_cost_keeps_custom_priced_slash_router_id(_local_model_cost_m } ) - selected = _select_model_name_for_cost_calc( + selected = select_model_name_for_cost_calc( model="vertex_ai/claude-opus-5", completion_response=None, custom_pricing=True, @@ -4358,7 +4358,7 @@ def test_explicit_pricing_precedes_private_provider_response_model( custom_pricing: bool, expected: str, ) -> None: - from litellm.cost_calculator import _select_model_name_for_cost_calc + from litellm.cost_calculator import select_model_name_for_cost_calc response = litellm.ModelResponse( id="x", @@ -4373,7 +4373,7 @@ def test_explicit_pricing_precedes_private_provider_response_model( ) response._hidden_params = {"provider_response_model": "selected-cost-model"} - selected = _select_model_name_for_cost_calc( + selected = select_model_name_for_cost_calc( model="requested-route", completion_response=response, base_model=base_model, diff --git a/tests/unit/test_deepseek_model_metadata.py b/tests/unit/test_deepseek_model_metadata.py index 91ed54b826c..769c28d3896 100644 --- a/tests/unit/test_deepseek_model_metadata.py +++ b/tests/unit/test_deepseek_model_metadata.py @@ -14,7 +14,7 @@ import os import litellm from litellm.utils import ( - _supports_factory, + supports_factory, ) # --------------------------------------------------------------------------- @@ -78,7 +78,7 @@ class TestBareModelFallback: # Simulate the pre-fix state: field missing from prefixed entry if key in litellm.model_cost: litellm.model_cost[key].pop("supports_response_schema", None) - result = _supports_factory( + result = supports_factory( model="deepseek-chat", custom_llm_provider="deepseek", key="supports_response_schema", @@ -100,7 +100,7 @@ class TestBareModelFallback: try: if key in litellm.model_cost: litellm.model_cost[key]["supports_function_calling"] = False - result = _supports_factory( + result = supports_factory( model="deepseek-reasoner", custom_llm_provider="deepseek", key="supports_function_calling", diff --git a/tests/unit/test_logging.py b/tests/unit/test_logging.py index da67cf60756..53a6555b652 100644 --- a/tests/unit/test_logging.py +++ b/tests/unit/test_logging.py @@ -34,7 +34,7 @@ from litellm._logging import ( _parse_json_logs_env, _plain_log_format, _stdout_truncation_marker, - _turn_on_json, + turn_on_json, format_base64_size, session_id_var, set_session_id, @@ -63,7 +63,7 @@ class CacheHitCustomLogger(CustomLogger): def test_json_mode_emits_one_record_per_logger(capfd): # Turn on JSON logging - _turn_on_json() + turn_on_json() # Make sure our loggers will emit INFO-level records for lg in (verbose_logger, verbose_router_logger, verbose_proxy_logger): lg.setLevel(logging.INFO) @@ -896,7 +896,7 @@ def test_disabled_diagnostic_call_does_not_render_arguments(caplog): def test_truncation_filter_survives_json_reconfiguration(): """The cap lives on the loggers, so swapping handlers (JSON mode) can't drop it.""" - _turn_on_json() + turn_on_json() for lg in (verbose_logger, verbose_router_logger, verbose_proxy_logger): assert any(isinstance(f, StdoutLogTruncationFilter) for f in lg.filters), f"{lg.name} lost stdout truncation" diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py index ffb17a17e3e..3f484592418 100644 --- a/tests/unit/test_main.py +++ b/tests/unit/test_main.py @@ -679,7 +679,7 @@ def test_build_database_url(): def test_bedrock_llama(): - litellm._turn_on_debug() + litellm.turn_on_debug() from litellm.types.utils import CallTypes from litellm.utils import return_raw_request @@ -847,7 +847,7 @@ def test_responses_api_bridge_check_strips_responses_prefix(): """Test that responses_api_bridge_check strips 'responses/' prefix and sets mode.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 4096} model_info, model = responses_api_bridge_check( @@ -881,7 +881,7 @@ def test_responses_api_bridge_check_gpt_5_4_tools_plus_reasoning_routes_to_respo """gpt-5.4 with both tools and reasoning_effort should route to Responses API.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.4", @@ -911,7 +911,7 @@ def test_responses_api_bridge_check_gpt_5_5_tools_plus_reasoning_routes_to_respo """gpt-5.5+ with both tools and reasoning_effort should route to Responses API.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.5-pro", @@ -928,7 +928,7 @@ def test_responses_api_bridge_check_azure_gpt_5_4_tools_plus_reasoning_routes_to """Azure gpt-5.4 with both tools and reasoning_effort should route to Responses API.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.4", @@ -949,7 +949,7 @@ def test_responses_api_bridge_check_azure_gpt_5_4_tools_with_default_reasoning_r """ from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.4", @@ -970,7 +970,7 @@ def test_responses_api_bridge_check_gpt_5_4_tools_with_default_reasoning_routes_ """ from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.4", @@ -1041,7 +1041,7 @@ def test_responses_api_bridge_check_gpt_5_6_tools_with_default_reasoning_routes_ monkeypatch.delenv("OPENAI_API_BASE", raising=False) monkeypatch.setattr(litellm, "api_base", None) - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model=model_name, @@ -1061,7 +1061,7 @@ def test_responses_api_bridge_check_gpt_5_4_tools_with_reasoning_none_stays_chat """ from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.4", @@ -1078,7 +1078,7 @@ def test_responses_api_bridge_check_reasoning_none_with_summary_still_routes_to_ """A reasoning summary is Responses-only regardless of effort value.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.4", @@ -1100,7 +1100,7 @@ def test_responses_api_bridge_check_gpt_5_4_custom_tools_only_stays_chat(): """ from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1117,7 +1117,7 @@ def test_responses_api_bridge_check_gpt_5_4_mixed_function_and_custom_tools_rout """One function tool in the mix is enough to make chat unservable with reasoning on.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1137,7 +1137,7 @@ def test_responses_api_bridge_check_gpt_5_4_flat_function_tool_routes_to_respons """Responses-style flat function tool defs still count as function tools.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1183,7 +1183,7 @@ def test_responses_api_bridge_check_dict_effort_none_stays_chat(): """The escape hatch must honor litellm's dict form: {"effort": "none"} means reasoning off.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1199,7 +1199,7 @@ def test_responses_api_bridge_check_dict_effort_none_stays_chat(): def test_responses_api_bridge_check_dict_effort_active_routes_to_responses(): from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1216,7 +1216,7 @@ def test_responses_api_bridge_check_dict_effort_none_with_summary_routes_to_resp """A summary inside the dict form is Responses-only even when effort is none.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1238,7 +1238,7 @@ def test_responses_api_bridge_check_blank_api_base_is_default_openai(blank_api_b """ from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1260,7 +1260,7 @@ def test_responses_api_bridge_check_custom_api_base_with_unset_effort_stays_chat """ from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1285,7 +1285,7 @@ def test_responses_api_bridge_check_custom_api_base_via_global_with_unset_effort from litellm.main import responses_api_bridge_check monkeypatch.setattr(litellm, "api_base", "http://vllm.internal:8000/v1") - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1312,7 +1312,7 @@ def test_responses_api_bridge_check_custom_api_base_via_env_with_unset_effort_st monkeypatch.delenv("OPENAI_BASE_URL", raising=False) monkeypatch.delenv("OPENAI_API_BASE", raising=False) monkeypatch.setenv(env_var, "http://vllm.internal:8000/v1") - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1406,7 +1406,7 @@ def test_responses_api_bridge_check_custom_api_base_with_explicit_effort_still_r """Explicit reasoning_effort keeps its pre-existing bridging behavior on any api_base.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.6", @@ -1424,7 +1424,7 @@ def test_responses_api_bridge_check_azure_with_api_base_and_unset_effort_routes( """Azure OpenAI always sets api_base and does enforce the constraint; keep bridging.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.4", @@ -1504,7 +1504,7 @@ def test_responses_api_bridge_check_older_gpt_5_tools_without_reasoning_stays_ch """Pre-5.4 GPT-5 names keep the old boundary: tools alone never bridge.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.1", @@ -1521,7 +1521,7 @@ def test_responses_api_bridge_check_gpt_5_4_reasoning_summary_without_tools_rout """gpt-5.4+ with reasoning_effort + reasoningSummary but no tools should bridge (AI SDK).""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5.4", @@ -1539,7 +1539,7 @@ def test_responses_api_bridge_check_gpt_5_reasoning_summary_routes_to_responses( """Bare ``gpt-5`` with reasoning_effort + reasoningSummary should bridge (not 5.4+).""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5", @@ -1557,7 +1557,7 @@ def test_responses_api_bridge_check_gpt_5_tools_without_summary_stays_chat(): """gpt-5 with tools + reasoning_effort but no summary should stay on chat.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 128000} model_info, model = responses_api_bridge_check( model="gpt-5", @@ -1891,7 +1891,7 @@ def test_responses_api_bridge_check_handles_exception(): """Test that responses_api_bridge_check handles exceptions and still processes responses/ models.""" from litellm.main import responses_api_bridge_check - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.side_effect = Exception("Model not found") model_info, model = responses_api_bridge_check( @@ -1921,7 +1921,7 @@ def test_responses_api_bridge_check_global_flag_does_not_affect_azure(): from litellm.main import responses_api_bridge_check with patch.object(litellm, "route_all_chat_openai_to_responses", True): - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 4096} model_info, model = responses_api_bridge_check( model="gpt-4o", @@ -1936,7 +1936,7 @@ def test_responses_api_bridge_check_global_flag_default_false(): from litellm.main import responses_api_bridge_check with patch.object(litellm, "route_all_chat_openai_to_responses", False): - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + with patch("litellm.main.get_model_info_helper") as mock_get_model_info: mock_get_model_info.return_value = {"max_tokens": 4096} model_info, model = responses_api_bridge_check( model="gpt-4o", @@ -4130,7 +4130,7 @@ def test_stream_chunk_builder_leaves_xai_reported_cost_to_the_calculator(monkeyp assert response is not None assert getattr(response.usage, "cost", None) == pytest.approx(0.42) assert response._hidden_params.get("response_cost") is None - assert logging_obj._response_cost_calculator(result=response) == pytest.approx(0.63) + assert logging_obj.response_cost_calculator(result=response) == pytest.approx(0.63) def test_speech_mistral_dispatches_and_decodes_audio(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): diff --git a/tests/unit/test_model_param_helper.py b/tests/unit/test_model_param_helper.py index a62779aeab1..b421cd6ef2a 100644 --- a/tests/unit/test_model_param_helper.py +++ b/tests/unit/test_model_param_helper.py @@ -3,8 +3,8 @@ from litellm.litellm_core_utils.model_param_helper import ModelParamHelper def test_cached_relevant_logging_args_matches_dynamic(): """Verify the cached frozenset matches the dynamically computed set.""" - cached = ModelParamHelper._relevant_logging_args - dynamic = ModelParamHelper._get_relevant_args_to_use_for_logging() + cached = ModelParamHelper.relevant_logging_args + dynamic = ModelParamHelper.get_relevant_args_to_use_for_logging() assert cached == dynamic assert isinstance(cached, frozenset) @@ -40,7 +40,7 @@ def test_get_all_llm_api_params_includes_responses_api(): otherwise Cache.get_cache_key() silently drops them and two requests that differ only in (e.g.) `instructions` collide on the same key. """ - all_params = ModelParamHelper._get_all_llm_api_params() + all_params = ModelParamHelper.get_all_llm_api_params() responses_only_params = { "instructions", "previous_response_id", diff --git a/tests/unit/test_private_usage_aliases.py b/tests/unit/test_private_usage_aliases.py new file mode 100644 index 00000000000..a3fa4ac27c3 --- /dev/null +++ b/tests/unit/test_private_usage_aliases.py @@ -0,0 +1,1052 @@ +from importlib import import_module +from types import MethodType +from typing import Final + +import pytest + +ALIAS_CASES: Final = ( + ( + "enterprise.enterprise_hooks.banned_keywords", + "", + "_ENTERPRISE_BannedKeywords", + "ENTERPRISE_BannedKeywords", + False, + ), + ( + "enterprise.enterprise_hooks.blocked_user_list", + "", + "_ENTERPRISE_BlockedUserList", + "ENTERPRISE_BlockedUserList", + False, + ), + ( + "enterprise.enterprise_hooks.google_text_moderation", + "", + "_ENTERPRISE_GoogleTextModeration", + "ENTERPRISE_GoogleTextModeration", + False, + ), + ( + "enterprise.enterprise_hooks.openai_moderation", + "", + "_ENTERPRISE_OpenAI_Moderation", + "ENTERPRISE_OpenAI_Moderation", + False, + ), + ( + "enterprise.litellm_enterprise.proxy.hooks.managed_files", + "", + "_PROXY_LiteLLMManagedFiles", + "PROXY_LiteLLMManagedFiles", + False, + ), + ( + "enterprise.litellm_enterprise.proxy.hooks.managed_vector_stores", + "", + "_PROXY_LiteLLMManagedVectorStores", + "PROXY_LiteLLMManagedVectorStores", + False, + ), + ("litellm", "", "_calculate_retry_after", "calculate_retry_after", False), + ("litellm", "", "_turn_on_debug", "turn_on_debug", False), + ("litellm", "", "_turn_on_json", "turn_on_json", False), + ("litellm._lazy_imports", "", "_get_default_encoding", "get_default_encoding", False), + ("litellm._lazy_imports", "", "_get_lazy_import_registry", "get_lazy_import_registry", False), + ("litellm._lazy_imports", "", "_get_messages_reach_token_count", "get_messages_reach_token_count", False), + ("litellm._lazy_imports", "", "_get_token_counter_new", "get_token_counter_new", False), + ("litellm._logging", "", "_is_debugging_on", "is_debugging_on", False), + ("litellm._logging", "", "_redact_string", "redact_string", False), + ("litellm._logging", "", "_turn_on_debug", "turn_on_debug", False), + ("litellm._logging", "", "_turn_on_json", "turn_on_json", False), + ("litellm._redis", "", "_generate_gcp_iam_access_token", "generate_gcp_iam_access_token", False), + ( + "litellm._redis_credential_provider", + "", + "_generate_gcp_iam_access_token", + "generate_gcp_iam_access_token", + False, + ), + ("litellm", "", "_should_retry", "should_retry", False), + ("litellm.batches.batch_utils", "", "_count_entry_tokens", "count_entry_tokens", False), + ("litellm.batches.batch_utils", "", "_estimate_batch_entry_tokens", "estimate_batch_entry_tokens", False), + ("litellm.batches.batch_utils", "", "_extract_file_access_credentials", "extract_file_access_credentials", False), + ("litellm.batches.batch_utils", "", "_get_file_content_as_dictionary", "get_file_content_as_dictionary", False), + ("litellm.batches.batch_utils", "", "_handle_completed_batch", "handle_completed_batch", False), + ("litellm.batches.batch_utils", "", "_iter_batch_input_lines", "iter_batch_input_lines", False), + ("litellm.caching.caching", "Cache", "_get_cache_logic", "get_cache_logic", False), + ( + "litellm.caching.caching", + "Cache", + "_get_preset_cache_key_from_kwargs", + "get_preset_cache_key_from_kwargs", + False, + ), + ("litellm.caching.caching", "Cache", "_supports_async", "supports_async", False), + ( + "litellm.caching.caching_handler", + "LLMCachingHandler", + "_add_streaming_response_to_cache", + "add_streaming_response_to_cache", + False, + ), + ("litellm.caching.caching_handler", "LLMCachingHandler", "_async_get_cache", "async_get_cache", False), + ( + "litellm.caching.caching_handler", + "LLMCachingHandler", + "_combine_cached_embedding_response_with_api_result", + "combine_cached_embedding_response_with_api_result", + False, + ), + ( + "litellm.caching.caching_handler", + "LLMCachingHandler", + "_should_store_result_in_cache", + "should_store_result_in_cache", + False, + ), + ( + "litellm.caching.caching_handler", + "LLMCachingHandler", + "_sync_add_streaming_response_to_cache", + "sync_add_streaming_response_to_cache", + False, + ), + ("litellm.caching.caching_handler", "LLMCachingHandler", "_sync_get_cache", "sync_get_cache", False), + ( + "litellm.completion_extras.litellm_responses_transformation.transformation", + "LiteLLMResponsesTransformationHandler", + "_map_reasoning_effort", + "map_reasoning_effort", + False, + ), + ("litellm.cost_calculator", "", "_infer_call_type", "infer_call_type", False), + ("litellm.cost_calculator", "", "_select_model_name_for_cost_calc", "select_model_name_for_cost_calc", False), + ( + "litellm.integrations.SlackAlerting.utils", + "", + "_add_langfuse_trace_id_to_alert", + "add_langfuse_trace_id_to_alert", + False, + ), + ( + "litellm.integrations.arize.arize_phoenix_prompt_manager", + "ArizePhoenixTemplateManager", + "_load_prompt_from_arize", + "load_prompt_from_arize", + False, + ), + ( + "litellm.integrations.bitbucket.bitbucket_prompt_manager", + "BitBucketTemplateManager", + "_load_prompt_from_bitbucket", + "load_prompt_from_bitbucket", + False, + ), + ( + "litellm.integrations.dotprompt", + "", + "_get_prompt_data_from_dotprompt_content", + "get_prompt_data_from_dotprompt_content", + False, + ), + ( + "litellm.integrations.dotprompt.prompt_manager", + "PromptManager", + "_parse_frontmatter", + "parse_frontmatter", + False, + ), + ("litellm.integrations.dynamodb", "DyanmoDBLogger", "_async_log_event", "async_log_event", False), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_build_filename", "build_filename", False), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_count_unique", "count_unique", False), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_database", "database", True), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_destination", "destination", True), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_serializer", "serializer", True), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_sum_column", "sum_column", False), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_transformer", "transformer", True), + ( + "litellm.integrations.gitlab.gitlab_prompt_manager", + "GitLabTemplateManager", + "_id_to_repo_path", + "id_to_repo_path", + False, + ), + ( + "litellm.integrations.gitlab.gitlab_prompt_manager", + "GitLabTemplateManager", + "_load_prompt_from_gitlab", + "load_prompt_from_gitlab", + False, + ), + ( + "litellm.integrations.opentelemetry", + "", + "_build_metric_attribute_filter", + "build_metric_attribute_filter", + False, + ), + ( + "litellm.integrations.opentelemetry", + "", + "_resolve_metric_attribute_filter", + "resolve_metric_attribute_filter", + False, + ), + ( + "litellm.integrations.prometheus_helpers", + "PrometheusLabelFactoryContext", + "_custom_by_sanitized_key", + "custom_by_sanitized_key", + True, + ), + ( + "litellm.integrations.prometheus_helpers", + "PrometheusLabelFactoryContext", + "_sanitized_enum", + "sanitized_enum", + True, + ), + ("litellm.integrations.prometheus_helpers", "PrometheusLabelFactoryContext", "_tag_labels", "tag_labels", True), + ( + "litellm.integrations.prometheus_helpers", + "", + "_get_cached_end_user_id_for_cost_tracking", + "get_cached_end_user_id_for_cost_tracking", + False, + ), + ( + "litellm.integrations.weave.weave_otel", + "", + "_get_weave_authorization_header", + "get_weave_authorization_header", + False, + ), + ( + "litellm.litellm_core_utils.core_helpers", + "", + "_get_parent_otel_span_from_kwargs", + "get_parent_otel_span_from_kwargs", + False, + ), + ("litellm.litellm_core_utils.dd_tracing", "", "_should_use_dd_profiler", "should_use_dd_profiler", False), + ( + "litellm.litellm_core_utils.exception_mapping_utils", + "", + "_add_key_name_and_team_to_alert", + "add_key_name_and_team_to_alert", + False, + ), + ( + "litellm.litellm_core_utils.get_litellm_params", + "", + "_get_base_model_from_litellm_call_metadata", + "get_base_model_from_litellm_call_metadata", + False, + ), + ( + "litellm.litellm_core_utils.get_llm_provider_logic", + "", + "_is_non_openai_azure_model", + "is_non_openai_azure_model", + False, + ), + ( + "litellm.litellm_core_utils.get_model_cost_map", + "GetModelCostMap", + "_get_backup_model_count", + "get_backup_model_count", + False, + ), + ( + "litellm.litellm_core_utils.health_check_helpers", + "HealthCheckHelpers", + "_update_model_params_with_health_check_tracking_information", + "update_model_params_with_health_check_tracking_information", + False, + ), + ( + "litellm.litellm_core_utils.health_check_utils", + "", + "_create_health_check_response", + "create_health_check_response", + False, + ), + ("litellm.litellm_core_utils.litellm_logging", "Logging", "_defer_async_logging", "defer_async_logging", True), + ( + "litellm.litellm_core_utils.litellm_logging", + "Logging", + "_enqueue_deferred_logging", + "enqueue_deferred_logging", + True, + ), + ("litellm.litellm_core_utils.litellm_logging", "Logging", "_get_trace_id", "get_trace_id", False), + ( + "litellm.litellm_core_utils.litellm_logging", + "Logging", + "_is_sync_litellm_request", + "is_sync_litellm_request", + False, + ), + ("litellm.litellm_core_utils.litellm_logging", "Logging", "_llm_caching_handler", "llm_caching_handler", True), + ( + "litellm.litellm_core_utils.litellm_logging", + "Logging", + "_on_detached_stream_failure", + "on_detached_stream_failure", + True, + ), + ( + "litellm.litellm_core_utils.litellm_logging", + "Logging", + "_response_cost_calculator", + "response_cost_calculator", + False, + ), + ( + "litellm.litellm_core_utils.litellm_logging", + "Logging", + "_update_completion_start_time", + "update_completion_start_time", + False, + ), + ( + "litellm.litellm_core_utils.litellm_logging", + "StandardLoggingPayloadSetup", + "_get_request_tags", + "get_request_tags", + False, + ), + ("litellm.litellm_core_utils.litellm_logging", "", "_get_masked_values", "get_masked_values", False), + ( + "litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking", + "StandardBuiltInToolCostTracking", + "_get_file_search_tool_call", + "get_file_search_tool_call", + False, + ), + ( + "litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking", + "StandardBuiltInToolCostTracking", + "_get_web_search_options", + "get_web_search_options", + False, + ), + ( + "litellm.litellm_core_utils.llm_cost_calc.utils", + "CostCalculatorUtils", + "_call_type_has_image_response", + "call_type_has_image_response", + False, + ), + ( + "litellm.litellm_core_utils.llm_cost_calc.utils", + "", + "_generic_cost_per_character", + "generic_cost_per_character", + False, + ), + ("litellm.litellm_core_utils.llm_cost_calc.utils", "", "_get_cost_per_unit", "get_cost_per_unit", False), + ( + "litellm.litellm_core_utils.llm_cost_calc.utils", + "", + "_get_regional_uplift_multiplier", + "get_regional_uplift_multiplier", + False, + ), + ( + "litellm.litellm_core_utils.llm_cost_calc.utils", + "", + "_get_service_tier_cost_key", + "get_service_tier_cost_key", + False, + ), + ("litellm.litellm_core_utils.llm_cost_calc.utils", "", "_is_above_128k", "is_above_128k", False), + ( + "litellm.litellm_core_utils.llm_request_utils", + "", + "_ensure_extra_body_is_safe", + "ensure_extra_body_is_safe", + False, + ), + ( + "litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response", + "", + "_handle_invalid_parallel_tool_calls", + "handle_invalid_parallel_tool_calls", + False, + ), + ( + "litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response", + "", + "_safe_convert_created_field", + "safe_convert_created_field", + False, + ), + ( + "litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response", + "", + "_should_convert_tool_call_to_json_mode", + "should_convert_tool_call_to_json_mode", + False, + ), + ( + "litellm.litellm_core_utils.logging_callback_manager", + "LoggingCallbackManager", + "_add_custom_callback_generic_api_str", + "add_custom_callback_generic_api_str", + False, + ), + ( + "litellm.litellm_core_utils.logging_callback_manager", + "LoggingCallbackManager", + "_get_all_callbacks", + "get_all_callbacks", + False, + ), + ( + "litellm.litellm_core_utils.logging_utils", + "", + "_assemble_complete_response_from_streaming_chunks", + "assemble_complete_response_from_streaming_chunks", + False, + ), + ( + "litellm.litellm_core_utils.model_param_helper", + "ModelParamHelper", + "_get_all_llm_api_params", + "get_all_llm_api_params", + False, + ), + ( + "litellm.litellm_core_utils.model_param_helper", + "ModelParamHelper", + "_get_relevant_args_to_use_for_logging", + "get_relevant_args_to_use_for_logging", + False, + ), + ( + "litellm.litellm_core_utils.model_param_helper", + "ModelParamHelper", + "_relevant_logging_args", + "relevant_logging_args", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.common_utils", + "", + "_audio_or_image_in_message_content", + "audio_or_image_in_message_content", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.common_utils", + "", + "_extract_reasoning_content", + "extract_reasoning_content", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.common_utils", + "", + "_get_image_mime_type_from_url", + "get_image_mime_type_from_url", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.common_utils", + "", + "_parse_content_for_reasoning", + "parse_content_for_reasoning", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.factory", + "BedrockConverseMessagesProcessor", + "_process_file_message", + "process_file_message", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.factory", + "BedrockImageProcessor", + "_validate_format", + "validate_format", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.factory", + "", + "_encode_tool_call_id_with_signature", + "encode_tool_call_id_with_signature", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.factory", + "", + "_get_thought_signature_from_tool", + "get_thought_signature_from_tool", + False, + ), + ("litellm.litellm_core_utils.prompt_templates.factory", "", "_parse_mime_type", "parse_mime_type", False), + ( + "litellm.litellm_core_utils.prompt_templates.huggingface_template_handler", + "", + "_aget_chat_template_file", + "aget_chat_template_file", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.huggingface_template_handler", + "", + "_aget_tokenizer_config", + "aget_tokenizer_config", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.huggingface_template_handler", + "", + "_extract_token_value", + "extract_token_value", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.huggingface_template_handler", + "", + "_get_chat_template_file", + "get_chat_template_file", + False, + ), + ( + "litellm.litellm_core_utils.prompt_templates.huggingface_template_handler", + "", + "_get_tokenizer_config", + "get_tokenizer_config", + False, + ), + ( + "litellm.litellm_core_utils.realtime_streaming", + "RealTimeStreaming", + "_detect_beta_header", + "detect_beta_header", + False, + ), + ( + "litellm.litellm_core_utils.realtime_streaming", + "RealTimeStreaming", + "_maybe_inject_guardrail_auto_response_disable", + "maybe_inject_guardrail_auto_response_disable", + False, + ), + ( + "litellm.litellm_core_utils.realtime_streaming", + "RealTimeStreaming", + "_session_created_sent_to_client", + "session_created_sent_to_client", + True, + ), + ("litellm.litellm_core_utils.secret_redaction", "", "_python_redact_string", "python_redact_string", False), + ( + "litellm.litellm_core_utils.secret_redaction", + "", + "_python_redact_structured_value", + "python_redact_structured_value", + False, + ), + ("litellm.litellm_core_utils.sensitive_data_masker", "SensitiveDataMasker", "_mask_value", "mask_value", False), + ( + "litellm.litellm_core_utils.streaming_handler", + "CustomStreamWrapper", + "_strip_sse_data_from_chunk", + "strip_sse_data_from_chunk", + False, + ), + ( + "litellm.proxy.enterprise.litellm_enterprise.proxy.utils", + "", + "_should_block_robots", + "should_block_robots", + False, + ), + ("litellm.realtime_api.main", "", "_realtime_health_check", "realtime_health_check", False), + ( + "litellm.responses.litellm_completion_transformation.transformation", + "LiteLLMCompletionResponsesConfig", + "_tool_call_id_from_responses_item", + "tool_call_id_from_responses_item", + False, + ), + ( + "litellm.responses.litellm_completion_transformation.transformation", + "LiteLLMCompletionResponsesConfig", + "_transform_chat_completion_annotations_to_response_output_annotations", + "transform_chat_completion_annotations_to_response_output_annotations", + False, + ), + ( + "litellm.responses.litellm_completion_transformation.transformation", + "LiteLLMCompletionResponsesConfig", + "_transform_chat_completion_usage_to_responses_usage", + "transform_chat_completion_usage_to_responses_usage", + False, + ), + ( + "litellm.responses.litellm_completion_transformation.transformation", + "LiteLLMCompletionResponsesConfig", + "_transform_tool_choice_for_responses_api_response", + "transform_tool_choice_for_responses_api_response", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_add_mcp_output_elements_to_response", + "add_mcp_output_elements_to_response", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_create_follow_up_input", + "create_follow_up_input", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_create_follow_up_messages_for_chat", + "create_follow_up_messages_for_chat", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_create_tool_execution_events", + "create_tool_execution_events", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_execute_tool_calls", + "execute_tool_calls", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_extract_tool_call_details", + "extract_tool_call_details", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_extract_tool_calls_from_chat_response", + "extract_tool_calls_from_chat_response", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_extract_tool_calls_from_response", + "extract_tool_calls_from_response", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_get_parent_request_tags", + "get_parent_request_tags", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_is_persistence_disabled", + "is_persistence_disabled", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_make_follow_up_call", + "make_follow_up_call", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_prepare_follow_up_call_params", + "prepare_follow_up_call_params", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_prepare_initial_call_params", + "prepare_initial_call_params", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_process_mcp_tools_to_openai_format", + "process_mcp_tools_to_openai_format", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_process_mcp_tools_without_openai_transform", + "process_mcp_tools_without_openai_transform", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_should_auto_execute_tools", + "should_auto_execute_tools", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_should_use_litellm_mcp_gateway", + "should_use_litellm_mcp_gateway", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_split_mcp_tools", + "split_mcp_tools", + False, + ), + ( + "litellm.responses.mcp.litellm_proxy_mcp_handler", + "LiteLLM_Proxy_MCP_Handler", + "_transform_mcp_tools_to_openai", + "transform_mcp_tools_to_openai", + False, + ), + ("litellm.responses.streaming_iterator", "", "_get_openai_response_types", "get_openai_response_types", False), + ("litellm.responses.utils", "ResponseAPILoggingUtils", "_is_response_api_usage", "is_response_api_usage", False), + ( + "litellm.responses.utils", + "ResponseAPILoggingUtils", + "_transform_response_api_usage_to_chat_usage", + "transform_response_api_usage_to_chat_usage", + False, + ), + ("litellm.responses.utils", "ResponsesAPIRequestUtils", "_build_container_id", "build_container_id", False), + ("litellm.responses.utils", "ResponsesAPIRequestUtils", "_decode_container_id", "decode_container_id", False), + ( + "litellm.responses.utils", + "ResponsesAPIRequestUtils", + "_decode_encrypted_item_id", + "decode_encrypted_item_id", + False, + ), + ( + "litellm.responses.utils", + "ResponsesAPIRequestUtils", + "_decode_responses_api_response_id", + "decode_responses_api_response_id", + False, + ), + ( + "litellm.responses.utils", + "ResponsesAPIRequestUtils", + "_encode_container_id_on_output_item", + "encode_container_id_on_output_item", + False, + ), + ( + "litellm.responses.utils", + "ResponsesAPIRequestUtils", + "_encode_container_ids_in_annotations", + "encode_container_ids_in_annotations", + False, + ), + ( + "litellm.responses.utils", + "ResponsesAPIRequestUtils", + "_restore_encrypted_content_item_ids_in_input", + "restore_encrypted_content_item_ids_in_input", + False, + ), + ( + "litellm.responses.utils", + "ResponsesAPIRequestUtils", + "_unwrap_encrypted_content_with_model_id", + "unwrap_encrypted_content_with_model_id", + False, + ), + ( + "litellm.responses.utils", + "ResponsesAPIRequestUtils", + "_update_responses_api_response_id_with_model_id", + "update_responses_api_response_id_with_model_id", + False, + ), + ( + "litellm.responses.utils", + "ResponsesAPIRequestUtils", + "_wrap_encrypted_content_with_model_id", + "wrap_encrypted_content_with_model_id", + False, + ), + ("litellm.router", "Router", "_are_all_deployments_blocked", "are_all_deployments_blocked", False), + ("litellm.router", "Router", "_deployment_usable_by_team", "deployment_usable_by_team", False), + ("litellm.router", "Router", "_get_all_deployments", "get_all_deployments", False), + ("litellm.router", "Router", "_get_model_from_alias", "get_model_from_alias", False), + ( + "litellm.router", + "Router", + "_get_router_deployment_budget_limiter", + "get_router_deployment_budget_limiter", + False, + ), + ("litellm.router", "Router", "_is_deployment_blocked", "is_deployment_blocked", False), + ( + "litellm.router", + "Router", + "_is_model_access_group_for_wildcard_route", + "is_model_access_group_for_wildcard_route", + False, + ), + ("litellm.router", "Router", "_replay_model_cost_registrations", "replay_model_cost_registrations", False), + ("litellm.router", "Router", "_routing_groups", "routing_groups", True), + ("litellm.router", "Router", "_update_redis_cache", "update_redis_cache", False), + ( + "litellm.router_strategy.adaptive_router.adaptive_router", + "AdaptiveRouter", + "_state_loaded", + "state_loaded", + True, + ), + ( + "litellm.router_strategy.tag_based_routing", + "", + "_get_tags_from_request_kwargs", + "get_tags_from_request_kwargs", + False, + ), + ("litellm.router_utils.add_retry_fallback_headers", "HiddenParamsAsyncIteratorWrapper", "_inner", "inner", True), + ("litellm.router_utils.add_retry_fallback_headers", "", "_HiddenParamsHost", "HiddenParamsHost", False), + ( + "litellm.router_utils.batch_utils", + "", + "_get_router_metadata_variable_name", + "get_router_metadata_variable_name", + False, + ), + ("litellm.router_utils.common_utils", "", "_is_proxy_admin_request", "is_proxy_admin_request", False), + ( + "litellm.router_utils.cooldown_callbacks", + "", + "_get_prometheus_logger_from_callbacks", + "get_prometheus_logger_from_callbacks", + False, + ), + ( + "litellm.router_utils.cooldown_handlers", + "", + "_async_get_cooldown_deployments", + "async_get_cooldown_deployments", + False, + ), + ( + "litellm.router_utils.cooldown_handlers", + "", + "_async_get_cooldown_deployments_with_debug_info", + "async_get_cooldown_deployments_with_debug_info", + False, + ), + ("litellm.router_utils.cooldown_handlers", "", "_get_cooldown_deployments", "get_cooldown_deployments", False), + ("litellm.router_utils.cooldown_handlers", "", "_set_cooldown_deployments", "set_cooldown_deployments", False), + ( + "litellm.router_utils.fallback_event_handlers", + "", + "_check_non_standard_fallback_format", + "check_non_standard_fallback_format", + False, + ), + ( + "litellm.router_utils.pattern_match_deployments", + "PatternMatchRouter", + "_pattern_to_regex", + "pattern_to_regex", + False, + ), + ("litellm.types.agents", "", "_normalize_a2a_jsonrpc_response", "normalize_a2a_jsonrpc_response", False), + ("litellm.types.completion", "", "_CompletionDispatchContext", "CompletionDispatchContext", False), + ( + "litellm.types.integrations.prometheus", + "", + "_sanitize_prometheus_label_name", + "sanitize_prometheus_label_name", + False, + ), + ( + "litellm.types.integrations.prometheus", + "", + "_sanitize_prometheus_label_value", + "sanitize_prometheus_label_value", + False, + ), + ("litellm.types.utils", "", "_generate_id", "generate_id", False), + ("litellm.utils", "ProviderConfigManager", "_get_bedrock_mantle_config", "get_bedrock_mantle_config", False), + ( + "litellm.utils", + "", + "_add_custom_logger_callback_to_specific_event", + "add_custom_logger_callback_to_specific_event", + False, + ), + ("litellm.utils", "", "_add_path_to_api_base", "add_path_to_api_base", False), + ("litellm.utils", "", "_apply_openai_param_overrides", "apply_openai_param_overrides", False), + ("litellm.utils", "", "_cached_get_model_info_helper", "cached_get_model_info_helper", False), + ("litellm.utils", "", "_calculate_retry_after", "calculate_retry_after", False), + ("litellm.utils", "", "_count_characters", "count_characters", False), + ("litellm.utils", "", "_get_base_model_from_metadata", "get_base_model_from_metadata", False), + ("litellm.utils", "", "_get_bundled_model_cost_map", "get_bundled_model_cost_map", False), + ("litellm.utils", "", "_get_deployment_order", "get_deployment_order", False), + ("litellm.utils", "", "_get_model_cost_key", "get_model_cost_key", False), + ("litellm.utils", "", "_get_model_info_helper", "get_model_info_helper", False), + ("litellm.utils", "", "_get_potential_model_names", "get_potential_model_names", False), + ("litellm.utils", "", "_remove_additional_properties", "remove_additional_properties", False), + ("litellm.utils", "", "_remove_json_schema_refs", "remove_json_schema_refs", False), + ("litellm.utils", "", "_remove_strict_from_schema", "remove_strict_from_schema", False), + ("litellm.utils", "", "_select_tokenizer", "select_tokenizer", False), + ("litellm.utils", "", "_should_retry", "should_retry", False), + ("litellm.utils", "", "_supports_factory", "supports_factory", False), + ("litellm.utils", "", "_update_dictionary", "update_dictionary", False), + ( + "litellm.vector_stores.vector_store_registry", + "VectorStoreIndexRegistry", + "_get_vector_store_indexes_from_db", + "get_vector_store_indexes_from_db", + False, + ), + ( + "litellm.vector_stores.vector_store_registry", + "VectorStoreRegistry", + "_get_vector_stores_from_db", + "get_vector_stores_from_db", + False, + ), +) + +PROPERTY_CASES: Final = ( + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_database", "database"), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_destination", "destination"), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_serializer", "serializer"), + ("litellm.integrations.focus.export_engine", "FocusExportEngine", "_transformer", "transformer"), + ( + "litellm.integrations.prometheus_helpers", + "PrometheusLabelFactoryContext", + "_custom_by_sanitized_key", + "custom_by_sanitized_key", + ), + ("litellm.integrations.prometheus_helpers", "PrometheusLabelFactoryContext", "_sanitized_enum", "sanitized_enum"), + ("litellm.integrations.prometheus_helpers", "PrometheusLabelFactoryContext", "_tag_labels", "tag_labels"), + ("litellm.litellm_core_utils.litellm_logging", "Logging", "_defer_async_logging", "defer_async_logging"), + ("litellm.litellm_core_utils.litellm_logging", "Logging", "_enqueue_deferred_logging", "enqueue_deferred_logging"), + ("litellm.litellm_core_utils.litellm_logging", "Logging", "_llm_caching_handler", "llm_caching_handler"), + ( + "litellm.litellm_core_utils.litellm_logging", + "Logging", + "_on_detached_stream_failure", + "on_detached_stream_failure", + ), + ( + "litellm.litellm_core_utils.realtime_streaming", + "RealTimeStreaming", + "_session_created_sent_to_client", + "session_created_sent_to_client", + ), + ("litellm.router", "Router", "_routing_groups", "routing_groups"), + ("litellm.router_strategy.adaptive_router.adaptive_router", "AdaptiveRouter", "_state_loaded", "state_loaded"), + ("litellm.router_utils.add_retry_fallback_headers", "HiddenParamsAsyncIteratorWrapper", "_inner", "inner"), +) + +CLASS_PROPERTY_CASES: Final = () + + +def _get_owner(module_path: str, owner_name: str) -> object: + module: Final = import_module(module_path or "litellm") + return module if not owner_name else getattr(module, owner_name) + + +def _get_instance(owner: object, public_name: str) -> object: + if not isinstance(owner, type): + raise TypeError(f"expected a class owner, got {type(owner).__name__}") + instance: Final = object.__new__(owner) + setattr(instance, public_name, object()) + return instance + + +@pytest.mark.parametrize( + ("module_path", "owner_name", "old_name", "new_name", "use_instance"), + ALIAS_CASES, +) +def test_public_aliases( + module_path: str, + owner_name: str, + old_name: str, + new_name: str, + use_instance: bool, +) -> None: + resolved_owner: Final = _get_owner(module_path, owner_name) + alias_owner: Final = _get_instance(resolved_owner, new_name) if use_instance else resolved_owner + old_value, new_value = getattr(alias_owner, old_name), getattr(alias_owner, new_name) + if isinstance(old_value, MethodType) and isinstance(new_value, MethodType): + assert old_value.__func__ is new_value.__func__ + else: + assert old_value is new_value + + +@pytest.mark.parametrize( + ("module_path", "owner_name", "old_name", "new_name"), + PROPERTY_CASES, +) +def test_protected_data_properties_round_trip( + module_path: str, + owner_name: str, + old_name: str, + new_name: str, +) -> None: + owner: Final = _get_owner(module_path, owner_name) + if not isinstance(owner, type): + raise TypeError(f"expected a class owner, got {type(owner).__name__}") + instance: Final = object.__new__(owner) + first_value: Final = object() + second_value: Final = object() + setattr(instance, new_name, first_value) + assert getattr(instance, old_name) is first_value + setattr(instance, old_name, second_value) + assert getattr(instance, new_name) is second_value + + +@pytest.mark.parametrize( + ("module_path", "owner_name", "old_name", "new_name"), + CLASS_PROPERTY_CASES, +) +def test_class_level_property_alias_round_trip( + module_path: str, + owner_name: str, + old_name: str, + new_name: str, +) -> None: + owner: Final = _get_owner(module_path, owner_name) + first_value: Final = object() + second_value: Final = object() + original_class_value: Final = getattr(owner, new_name) + try: + setattr(owner, old_name, first_value) + assert getattr(owner, new_name) is first_value + setattr(owner, new_name, second_value) + assert getattr(owner, old_name) is second_value + finally: + setattr(owner, new_name, original_class_value) diff --git a/tests/unit/test_redact_string_in_error_paths.py b/tests/unit/test_redact_string_in_error_paths.py index 74da68208cb..eaf2d43a25f 100644 --- a/tests/unit/test_redact_string_in_error_paths.py +++ b/tests/unit/test_redact_string_in_error_paths.py @@ -1,5 +1,5 @@ """ -Tests for _redact_string usage in error/logging paths. +Tests for redact_string usage in error/logging paths. Covers actual execution of redaction in: - WebSocket close reasons in realtime handlers (openai, bedrock) @@ -15,36 +15,36 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from litellm._logging import _ENABLE_SECRET_REDACTION, _redact_string +from litellm._logging import _ENABLE_SECRET_REDACTION, redact_string class TestRedactStringFunction: def test_redacts_bearer_token(self): text = "Authorization: Bearer sk-9876567890abcdefghij" - result = _redact_string(text) + result = redact_string(text) assert "sk-9876567890abcdefghij" not in result assert "REDACTED" in result def test_redacts_api_key_in_url(self): text = "Error at https://example.com?api_key=my-secret-key-value-here" - result = _redact_string(text) + result = redact_string(text) assert "my-secret-key-value-here" not in result def test_redacts_google_api_key(self): text = "key=AIzaSyB1234567890abcdefghijklmnopqrstuvwx" - result = _redact_string(text) + result = redact_string(text) assert "AIzaSyB1234567890abcdefghijklmnopqrstuvwx" not in result def test_passes_clean_text_through(self): text = "This is a normal error message with no secrets" - assert _redact_string(text) == text + assert redact_string(text) == text @pytest.mark.skipif( not _ENABLE_SECRET_REDACTION, reason="redaction disabled via env var" ) def test_redaction_enabled_by_default(self): text = "Bearer sk-9876567890abcdefghij" - result = _redact_string(text) + result = redact_string(text) assert "sk-9876567890abcdefghij" not in result @@ -93,26 +93,26 @@ class TestOpenAIRealtimeRedaction: class TestBedrockRealtimeRedaction: - """Test that _redact_string produces safe close reasons for Bedrock-style errors.""" + """Test that redact_string produces safe close reasons for Bedrock-style errors.""" def test_internal_error_message_redacted(self): secret_error = RuntimeError( "Failed with aws_secret_access_key=AKIAIOSFODNN7EXAMPLE123456" ) - reason = _redact_string(f"Internal error: {str(secret_error)}") + reason = redact_string(f"Internal error: {str(secret_error)}") assert "AKIAIOSFODNN7EXAMPLE123456" not in reason class TestLLMHTTPHandlerRealtimeRedaction: - """Test _redact_string on the exact patterns used in llm_http_handler WS close.""" + """Test redact_string on the exact patterns used in llm_http_handler WS close.""" def test_invalid_status_pattern(self): error_msg = "InvalidStatusCode: 403 for wss://api.example.com?api_key=sk-leaked-key-here" - assert "sk-leaked-key-here" not in _redact_string(str(error_msg)) + assert "sk-leaked-key-here" not in redact_string(str(error_msg)) def test_internal_server_error_pattern(self): error_msg = "Connection failed for api_key=sk-secret-key-12345678" - assert "sk-secret-key-12345678" not in _redact_string( + assert "sk-secret-key-12345678" not in redact_string( f"Internal server error: {error_msg}" ) @@ -126,7 +126,7 @@ class TestProxyStreamingDataGeneratorRedaction: except RuntimeError: raw_tb = traceback.format_exc() - redacted_tb = _redact_string(raw_tb) + redacted_tb = redact_string(raw_tb) assert "sk-9876567890abcdefghij" not in redacted_tb assert "Traceback" in redacted_tb diff --git a/tests/unit/test_redis.py b/tests/unit/test_redis.py index 301ae2573e4..7793ba977b4 100644 --- a/tests/unit/test_redis.py +++ b/tests/unit/test_redis.py @@ -1143,7 +1143,7 @@ def test_gcp_iam_credential_provider_get_credentials(): service_account = "projects/-/serviceAccounts/test@project.iam.gserviceaccount.com" with patch( - "litellm._redis_credential_provider._generate_gcp_iam_access_token", + "litellm._redis_credential_provider.generate_gcp_iam_access_token", return_value="tok-1", ) as mock_gen: provider = GCPIAMCredentialProvider(service_account) @@ -1161,7 +1161,7 @@ def test_gcp_iam_credential_provider_caches_token(): service_account = "projects/-/serviceAccounts/test@project.iam.gserviceaccount.com" with patch( - "litellm._redis_credential_provider._generate_gcp_iam_access_token", + "litellm._redis_credential_provider.generate_gcp_iam_access_token", return_value="tok-cached", ) as mock_gen: provider = GCPIAMCredentialProvider(service_account) @@ -1184,7 +1184,7 @@ def test_gcp_iam_credential_provider_refreshes_on_expiry(): service_account = "projects/-/serviceAccounts/test@project.iam.gserviceaccount.com" with patch( - "litellm._redis_credential_provider._generate_gcp_iam_access_token", + "litellm._redis_credential_provider.generate_gcp_iam_access_token", side_effect=["tok-1", "tok-2"], ) as mock_gen: provider = GCPIAMCredentialProvider(service_account) @@ -1210,7 +1210,7 @@ def test_gcp_iam_credential_provider_cache_shared_across_instances(): service_account = "projects/-/serviceAccounts/shared@project.iam.gserviceaccount.com" with patch( - "litellm._redis_credential_provider._generate_gcp_iam_access_token", + "litellm._redis_credential_provider.generate_gcp_iam_access_token", return_value="tok-shared", ) as mock_gen: p1 = GCPIAMCredentialProvider(service_account) diff --git a/tests/unit/test_register_model_custom_pricing.py b/tests/unit/test_register_model_custom_pricing.py index fa4fee8f6d8..1f6890940d1 100644 --- a/tests/unit/test_register_model_custom_pricing.py +++ b/tests/unit/test_register_model_custom_pricing.py @@ -801,23 +801,30 @@ def test_update_dictionary_merges_nested_dicts_without_aliasing(): object stays untouched, and the caller's incoming nested dict is never inserted by reference into the merged result. """ - from litellm.utils import _update_dictionary + from litellm.utils import update_dictionary - existing_nested = {"hours_utc": "01:00-02:00"} + from typing import cast + + existing_nested = cast(dict[str, object], {"hours_utc": "01:00-02:00", 1: "existing"}) existing = {"off_peak_pricing": existing_nested} - incoming_nested = {"windows": [{"hours_utc": "16:00-19:00", "weekdays": [2]}]} + incoming_nested = cast( + dict[str, object], + {"windows": [{"hours_utc": "16:00-19:00", "weekdays": [2]}], 2: "new"}, + ) incoming = {"off_peak_pricing": incoming_nested} - merged = _update_dictionary(existing, incoming) + merged = update_dictionary(existing, incoming) assert merged["off_peak_pricing"] == { "hours_utc": "01:00-02:00", "windows": [{"hours_utc": "16:00-19:00", "weekdays": [2]}], + 1: "existing", + 2: "new", } - assert existing_nested == {"hours_utc": "01:00-02:00"} + assert existing_nested == {"hours_utc": "01:00-02:00", 1: "existing"} assert merged["off_peak_pricing"] is not incoming_nested - fresh = _update_dictionary({}, incoming) + fresh = update_dictionary({}, incoming) assert fresh["off_peak_pricing"] == incoming_nested assert fresh["off_peak_pricing"] is not incoming_nested diff --git a/tests/unit/test_responses_api_bridge_non_stream.py b/tests/unit/test_responses_api_bridge_non_stream.py index 617b2cfc031..3b38913f011 100644 --- a/tests/unit/test_responses_api_bridge_non_stream.py +++ b/tests/unit/test_responses_api_bridge_non_stream.py @@ -138,7 +138,7 @@ def test_transform_usage_no_token_details(): ) # Transform to Responses API usage format - responses_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + responses_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( completion_response ) @@ -169,7 +169,7 @@ def test_transform_usage_with_cached_tokens_only(): reasoning_tokens=None, # No reasoning tokens ) - responses_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + responses_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( completion_response ) @@ -207,7 +207,7 @@ def test_transform_usage_maps_nested_cache_creation_input_tokens(): }, ) - responses_usage: Final = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + responses_usage: Final = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( usage ) @@ -230,7 +230,7 @@ def test_transform_usage_with_reasoning_tokens_only(): reasoning_tokens=60, # Has reasoning tokens ) - responses_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + responses_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( completion_response ) @@ -266,7 +266,7 @@ def test_transform_usage_with_both_token_details(): text_tokens=50, # Also include text_tokens ) - responses_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + responses_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( completion_response ) @@ -306,7 +306,7 @@ def test_transform_usage_with_zero_values(): reasoning_tokens=0, # Explicitly 0 — preserved ) - responses_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + responses_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( completion_response ) @@ -338,7 +338,7 @@ def test_transform_usage_unknown_reasoning_split_keeps_output_tokens_details(): completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=None, text_tokens=None), ) - responses_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage(usage) + responses_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage(usage) assert responses_usage.output_tokens_details is not None assert responses_usage.output_tokens_details.reasoning_tokens == 0 @@ -436,7 +436,7 @@ def test_all_providers_transformation_scenarios(): ) # This should not raise any errors - responses_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + responses_usage = LiteLLMCompletionResponsesConfig.transform_chat_completion_usage_to_responses_usage( completion_response ) diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 03b53f36e6f..08d529a39d3 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -175,31 +175,18 @@ def test_router_model_group_encrypted_content_affinity_callback_registration(): num_retries=0, ) callbacks = router.optional_callbacks or [] - encrypted_content_callbacks = [ - cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) - ] - deployment_callback = next( - cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck) - ) + encrypted_content_callbacks = [cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck)] + deployment_callback = next(cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck)) assert len(encrypted_content_callbacks) == 1 assert encrypted_content_callbacks[0].enable_global_affinity is False - assert ( - encrypted_content_callbacks[0].model_group_affinity_config - == model_group_affinity_config - ) - assert callbacks.index(encrypted_content_callbacks[0]) < callbacks.index( - deployment_callback - ) - assert litellm.callbacks.index(encrypted_content_callbacks[0]) < ( - litellm.callbacks.index(deployment_callback) - ) + assert encrypted_content_callbacks[0].model_group_affinity_config == model_group_affinity_config + assert callbacks.index(encrypted_content_callbacks[0]) < callbacks.index(deployment_callback) + assert litellm.callbacks.index(encrypted_content_callbacks[0]) < (litellm.callbacks.index(deployment_callback)) router._add_encrypted_content_affinity_check(enable_global_affinity=True) callbacks = router.optional_callbacks or [] - encrypted_content_callbacks = [ - cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) - ] + encrypted_content_callbacks = [cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck)] assert len(encrypted_content_callbacks) == 1 assert encrypted_content_callbacks[0].enable_global_affinity is True assert encrypted_content_callbacks[0].router is router @@ -230,13 +217,9 @@ async def test_encrypted_content_affinity_model_group_config_is_additive(): }, target_deployment, ] - encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id( - "deployment-b", "rs_test" - ) + encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-b", "rs_test") - assert EncryptedContentAffinityCheck.has_model_group_affinity_enabled( - {model_group: ["encrypted_content_affinity"]} - ) + assert EncryptedContentAffinityCheck.has_model_group_affinity_enabled({model_group: ["encrypted_content_affinity"]}) assert not EncryptedContentAffinityCheck.has_model_group_affinity_enabled(None) per_group_check = EncryptedContentAffinityCheck( @@ -277,10 +260,7 @@ async def test_encrypted_content_affinity_model_group_config_is_additive(): ) assert unfiltered == healthy_deployments - assert ( - "encrypted_content_affinity_enabled" - not in disabled_request_kwargs["litellm_metadata"] - ) + assert "encrypted_content_affinity_enabled" not in disabled_request_kwargs["litellm_metadata"] global_check = EncryptedContentAffinityCheck( enable_global_affinity=True, @@ -300,9 +280,7 @@ async def test_encrypted_content_affinity_model_group_config_is_additive(): ) assert globally_filtered == [target_deployment] - assert global_request_kwargs["litellm_metadata"][ - "encrypted_content_affinity_enabled" - ] + assert global_request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] @pytest.mark.asyncio @@ -349,18 +327,10 @@ async def test_encrypted_content_affinity_takes_priority_over_user_key_affinity( num_retries=0, ) callbacks = router.optional_callbacks or [] - deployment_callback = next( - cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck) - ) - encrypted_content_callback = next( - cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) - ) - assert callbacks.index(encrypted_content_callback) < callbacks.index( - deployment_callback - ) - assert litellm.callbacks.index(encrypted_content_callback) < ( - litellm.callbacks.index(deployment_callback) - ) + deployment_callback = next(cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck)) + encrypted_content_callback = next(cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck)) + assert callbacks.index(encrypted_content_callback) < callbacks.index(deployment_callback) + assert litellm.callbacks.index(encrypted_content_callback) < (litellm.callbacks.index(deployment_callback)) cache_key = DeploymentAffinityCheck.get_affinity_cache_key( model_group=model_group, @@ -371,9 +341,7 @@ async def test_encrypted_content_affinity_takes_priority_over_user_key_affinity( value={"model_id": "deployment-a"}, ttl=60, ) - encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id( - "deployment-b", "rs_test" - ) + encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-b", "rs_test") request_kwargs = { "input": [{"type": "reasoning", "id": encoded_id}], "litellm_metadata": {"user_api_key_hash": user_api_key_hash}, @@ -559,7 +527,9 @@ async def test_async_router_acreate_file_passthrough_keeps_the_file_and_forwards from io import BytesIO from unittest.mock import MagicMock, patch - jsonl_content = b'{"custom_id": "r1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "vertex-batch"}}\n' + jsonl_content = ( + b'{"custom_id": "r1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "vertex-batch"}}\n' + ) router = litellm.Router( model_list=[ { @@ -570,9 +540,7 @@ async def test_async_router_acreate_file_passthrough_keeps_the_file_and_forwards ) with patch("litellm.acreate_file", return_value=MagicMock()) as mock_acreate_file: - await router.acreate_file( - model="vertex-batch", purpose="batch", file=BytesIO(jsonl_content), passthrough=True - ) + await router.acreate_file(model="vertex-batch", purpose="batch", file=BytesIO(jsonl_content), passthrough=True) forwarded = mock_acreate_file.call_args.kwargs assert forwarded["passthrough"] is True forwarded["file"].seek(0) @@ -907,8 +875,6 @@ async def test_arouter_async_get_healthy_deployments(): assert result[0]["litellm_params"]["model"] == "gpt-3.5-turbo" - - def test_arouter_test_team_model(): """ Test that router.test_team_model returns the correct model @@ -1020,9 +986,7 @@ async def test_arouter_aretrieve_batch(): ], ) - with patch.object( - litellm, "aretrieve_batch", return_value=AsyncMock() - ) as mock_aretrieve_batch: + with patch.object(litellm, "aretrieve_batch", return_value=AsyncMock()) as mock_aretrieve_batch: try: response = await router.aretrieve_batch( model="gpt-3.5-turbo", @@ -1282,6 +1246,7 @@ def test_sync_deployment_callback_on_success_skips_batch_retrieves( == expected_successes ) + _ROUTING_STRATEGY_CACHE_MARKERS = ("_map", "_request_count", ":tpm:", ":rpm:") @@ -1293,8 +1258,7 @@ async def _moved_routing_counters(router, timeout: float = 2.0) -> list[str]: moved = sorted( f"{key}={cache_dict[key]}" for key in cache_dict - if any(marker in key for marker in _ROUTING_STRATEGY_CACHE_MARKERS) - and cache_dict[key] + if any(marker in key for marker in _ROUTING_STRATEGY_CACHE_MARKERS) and cache_dict[key] ) if moved: return moved @@ -1375,9 +1339,7 @@ async def test_arouter_aretrieve_file_content(): Test that router.acreate_file with JSONL file returns the correct response """ - with patch.object( - litellm, "afile_content", return_value=AsyncMock() - ) as mock_afile_content: + with patch.object(litellm, "afile_content", return_value=AsyncMock()) as mock_afile_content: router = litellm.Router( model_list=[ { @@ -1438,7 +1400,7 @@ async def test_arouter_filter_team_based_models(): assert result is not None # FAILS - with pytest.raises(Exception, match='No deployments available for selected model, Try again in') as e: + with pytest.raises(Exception, match="No deployments available for selected model, Try again in") as e: result = await router.acompletion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello, world!"}], @@ -1522,9 +1484,7 @@ def test_arouter_should_include_deployment(): model=deployment_with_team_and_public_name, team_id="test-team", ) - assert ( - result is True - ), "Should return True when team_id and team_public_model_name match" + assert result is True, "Should return True when team_id and team_public_model_name match" # Test Case 2: Team-specific deployment - team_id matches but model_name doesn't match team_public_model_name result = router.should_include_deployment( @@ -1532,9 +1492,9 @@ def test_arouter_should_include_deployment(): model=deployment_with_team_and_public_name, team_id="test-team", ) - assert ( - result is False - ), "Should return False when team_id matches but model_name doesn't match team_public_model_name" + assert result is False, ( + "Should return False when team_id matches but model_name doesn't match team_public_model_name" + ) # Test Case 3: Team-specific deployment - team_id doesn't match result = router.should_include_deployment( @@ -1550,30 +1510,18 @@ def test_arouter_should_include_deployment(): model=deployment_with_team_no_public_name, team_id="test-team", ) - assert ( - result is True - ), "Should return True when team deployment has no team_public_model_name to match" + assert result is True, "Should return True when team deployment has no team_public_model_name to match" # Test Case 5: Non-team deployment - model_name matches and no team_id - result = router.should_include_deployment( - model_name="gpt-4", model=deployment_without_team, team_id=None - ) - assert ( - result is True - ), "Should return True when model_name matches and deployment has no team_id" + result = router.should_include_deployment(model_name="gpt-4", model=deployment_without_team, team_id=None) + assert result is True, "Should return True when model_name matches and deployment has no team_id" # Test Case 6: Non-team deployment - model_name matches but team_id provided (should still work) - result = router.should_include_deployment( - model_name="gpt-4", model=deployment_without_team, team_id="any-team" - ) - assert ( - result is True - ), "Should return True when model_name matches non-team deployment, regardless of team_id param" + result = router.should_include_deployment(model_name="gpt-4", model=deployment_without_team, team_id="any-team") + assert result is True, "Should return True when model_name matches non-team deployment, regardless of team_id param" # Test Case 7: Non-team deployment - model_name doesn't match - result = router.should_include_deployment( - model_name="different-model", model=deployment_without_team, team_id=None - ) + result = router.should_include_deployment(model_name="different-model", model=deployment_without_team, team_id=None) assert result is False, "Should return False when model_name doesn't match" # Test Case 8: Team deployment accessed without matching team_id @@ -1582,9 +1530,7 @@ def test_arouter_should_include_deployment(): model=deployment_with_team_and_public_name, team_id=None, ) - assert ( - result is True - ), "Should return True when matching model with exact model_name" + assert result is True, "Should return True when matching model with exact model_name" def test_arouter_responses_api_bridge(): @@ -1634,9 +1580,7 @@ def test_arouter_responses_api_bridge(): "status": "completed", "output": [], } - mock_response.text = ( - '{"id": "resp_test", "object": "response", "status": "completed", "output": []}' - ) + mock_response.text = '{"id": "resp_test", "object": "response", "status": "completed", "output": []}' with patch.object(client, "post", return_value=mock_response) as mock_post: try: @@ -1710,7 +1654,7 @@ def test_add_invalid_provider_to_router(): ], ) - with pytest.raises(Exception, match='Unsupported provider - vertex_ai_eu') as e: + with pytest.raises(Exception, match="Unsupported provider - vertex_ai_eu") as e: router.add_deployment( Deployment( model_name="vertex_ai/*", @@ -1735,7 +1679,9 @@ def registered_custom_provider(monkeypatch: pytest.MonkeyPatch) -> str: model="gpt-5.6", messages=[{"role": "user", "content": "hi"}], mock_response="served by onprem handler" ) - monkeypatch.setattr(litellm, "custom_provider_map", [{"provider": "test-onprem-llm", "custom_handler": OnPremLLM()}]) + monkeypatch.setattr( + litellm, "custom_provider_map", [{"provider": "test-onprem-llm", "custom_handler": OnPremLLM()}] + ) monkeypatch.setattr(litellm, "provider_list", list(litellm.provider_list)) monkeypatch.setattr(litellm, "_custom_providers", list(litellm._custom_providers)) return "test-onprem-llm" @@ -1812,15 +1758,9 @@ async def test_router_ageneric_api_call_with_fallbacks_helper(): }, } - with patch.object( - router, "_update_kwargs_with_deployment" - ) as mock_update_kwargs: - with patch.object( - router, "async_routing_strategy_pre_call_checks" - ) as mock_pre_call_checks: - with patch.object( - router, "_get_client", return_value=None - ) as mock_get_client: + with patch.object(router, "_update_kwargs_with_deployment") as mock_update_kwargs: + with patch.object(router, "async_routing_strategy_pre_call_checks") as mock_pre_call_checks: + with patch.object(router, "_get_client", return_value=None) as mock_get_client: result = await router._ageneric_api_call_with_fallbacks_helper( model="gpt-3.5-turbo", original_generic_function=mock_generic_function, @@ -1855,7 +1795,7 @@ async def test_router_ageneric_api_call_with_fallbacks_helper(): with patch.object(router, "async_get_available_deployment") as mock_get_deployment: mock_get_deployment.side_effect = Exception("No deployment available") - with pytest.raises(Exception, match='No deployment available') as exc_info: + with pytest.raises(Exception, match="No deployment available") as exc_info: await router._ageneric_api_call_with_fallbacks_helper( model="gpt-3.5-turbo", original_generic_function=mock_generic_function, @@ -1885,15 +1825,9 @@ async def test_router_ageneric_api_call_with_fallbacks_helper(): max_parallel_requests=1, model_id="deployment-1", model_group="gpt-3.5-turbo" ) - with patch.object( - router, "_update_kwargs_with_deployment" - ) as mock_update_kwargs: - with patch.object( - router, "_get_client", return_value=mock_semaphore - ) as mock_get_client: - with patch.object( - router, "async_routing_strategy_pre_call_checks" - ) as mock_pre_call_checks: + with patch.object(router, "_update_kwargs_with_deployment") as mock_update_kwargs: + with patch.object(router, "_get_client", return_value=mock_semaphore) as mock_get_client: + with patch.object(router, "async_routing_strategy_pre_call_checks") as mock_pre_call_checks: result = await router._ageneric_api_call_with_fallbacks_helper( model="gpt-3.5-turbo", original_generic_function=mock_semaphore_function, @@ -1922,16 +1856,10 @@ async def test_router_ageneric_api_call_with_fallbacks_helper(): }, } - with patch.object( - router, "_update_kwargs_with_deployment" - ) as mock_update_kwargs: - with patch.object( - router, "_get_client", return_value=None - ) as mock_get_client: - with patch.object( - router, "async_routing_strategy_pre_call_checks" - ) as mock_pre_call_checks: - with pytest.raises(Exception, match='Mock failure') as exc_info: + with patch.object(router, "_update_kwargs_with_deployment") as mock_update_kwargs: + with patch.object(router, "_get_client", return_value=None) as mock_get_client: + with patch.object(router, "async_routing_strategy_pre_call_checks") as mock_pre_call_checks: + with pytest.raises(Exception, match="Mock failure") as exc_info: await router._ageneric_api_call_with_fallbacks_helper( model="gpt-3.5-turbo", original_generic_function=mock_failing_function, @@ -1999,9 +1927,9 @@ async def test_ageneric_api_call_deployment_model_overrides_alias(): original_generic_function=capture_model, ) - assert ( - captured["model"] == "vertex_ai/gemini-2.5-flash" - ), f"Expected deployment model 'vertex_ai/gemini-2.5-flash', got '{captured['model']}'" + assert captured["model"] == "vertex_ai/gemini-2.5-flash", ( + f"Expected deployment model 'vertex_ai/gemini-2.5-flash', got '{captured['model']}'" + ) @pytest.mark.asyncio @@ -2107,14 +2035,10 @@ def test_router_get_model_access_groups_team_only_models(): ] ) - access_groups = router.get_model_access_groups( - model_name="gpt-3.5-turbo", team_id=None - ) + access_groups = router.get_model_access_groups(model_name="gpt-3.5-turbo", team_id=None) assert len(access_groups) == 0 - access_groups = router.get_model_access_groups( - model_name="gpt-3.5-turbo", team_id="team_1" - ) + access_groups = router.get_model_access_groups(model_name="gpt-3.5-turbo", team_id="team_1") assert list(access_groups.keys()) == ["default-models"] @@ -2209,9 +2133,7 @@ def test_model_group_info_cost_from_db_model_info(): ] ) - with patch.object( - router, "get_deployment_model_info", side_effect=Exception("not found") - ): + with patch.object(router, "get_deployment_model_info", side_effect=Exception("not found")): result = router._cached_get_model_group_info("my-custom-model") assert result is not None assert result.input_cost_per_token == 0.0001 @@ -2239,9 +2161,7 @@ def test_model_group_info_cost_none_when_db_model_info_has_no_cost(): ] ) - with patch.object( - router, "get_deployment_model_info", side_effect=Exception("not found") - ): + with patch.object(router, "get_deployment_model_info", side_effect=Exception("not found")): result = router._cached_get_model_group_info("my-custom-model-no-cost") assert result is not None assert result.input_cost_per_token is None @@ -2535,9 +2455,7 @@ def test_model_group_info_with_stringified_cost_values(): } return None - with patch.object( - router, "get_deployment_model_info", side_effect=_model_info_with_str_costs - ): + with patch.object(router, "get_deployment_model_info", side_effect=_model_info_with_str_costs): result = router._set_model_group_info( model_group="my-custom-model", user_facing_model_group_name="my-custom-model", @@ -2583,9 +2501,7 @@ def test_model_group_info_db_fallback_with_stringified_cost_values(): ] ) - with patch.object( - router, "get_deployment_model_info", side_effect=Exception("not found") - ): + with patch.object(router, "get_deployment_model_info", side_effect=Exception("not found")): result = router._set_model_group_info( model_group="my-custom-model", user_facing_model_group_name="my-custom-model", @@ -2845,6 +2761,7 @@ async def test_acompletion_streaming_iterator(): # Collect streamed chunks — the first chunk succeeds, then the error re-raises collected_chunks = [] + async def _drain(): async for chunk in result: collected_chunks.append(chunk) @@ -3192,9 +3109,7 @@ def test_adopt_fallback_response_headers_replaces_rather_than_merges(): "additional_headers": {"llm_provider-x-request-id": "req-FALLBACK"}, } - wrapper.adopt_fallback_response_headers( - fallback, Router._prepare_fallback_hidden_params(fallback) - ) + wrapper.adopt_fallback_response_headers(fallback, Router._prepare_fallback_hidden_params(fallback)) assert wrapper._response_headers == {"x-request-id": "req-FALLBACK"} assert wrapper._hidden_params["model_id"] == "fallback-deployment" @@ -3266,9 +3181,7 @@ def test_adopt_fallback_response_headers_drops_headers_the_fallback_cannot_repla fallback._response_headers = None fallback._hidden_params = {"model_id": "fallback-deployment"} - wrapper.adopt_fallback_response_headers( - fallback, Router._prepare_fallback_hidden_params(fallback) - ) + wrapper.adopt_fallback_response_headers(fallback, Router._prepare_fallback_hidden_params(fallback)) assert wrapper._response_headers is None assert wrapper._hidden_params["model_id"] == "fallback-deployment" @@ -3295,9 +3208,7 @@ def test_adopt_fallback_response_headers_keeps_identity_when_fallback_has_none() hidden_params_before = wrapper._hidden_params fallback = object() - wrapper.adopt_fallback_response_headers( - fallback, Router._prepare_fallback_hidden_params(fallback) - ) + wrapper.adopt_fallback_response_headers(fallback, Router._prepare_fallback_hidden_params(fallback)) assert wrapper._response_headers is None assert wrapper._hidden_params is hidden_params_before @@ -3371,9 +3282,7 @@ async def test_set_response_headers_is_the_only_complexity_header_source_for_pro **additional_headers, ) - assert not { - key for key in proxy_headers if key.startswith("x-litellm-complexity-router-") - } + assert not {key for key in proxy_headers if key.startswith("x-litellm-complexity-router-")} @pytest.mark.asyncio @@ -4292,11 +4201,7 @@ def _make_responses_iterator( BaseResponsesAPIStreamingIterator, ) - base = ( - LiteLLMCompletionStreamingIterator - if bridge - else BaseResponsesAPIStreamingIterator - ) + base = LiteLLMCompletionStreamingIterator if bridge else BaseResponsesAPIStreamingIterator class _Iter(base): def __init__(self): @@ -4396,9 +4301,7 @@ async def test_aresponses_streaming_iterator_fallback(): BaseResponsesAPIStreamingIterator, ) - router = _make_router_with_fallback( - "anthropic/claude-sonnet-4-6", "vertex_ai/claude-sonnet-4-6" - ) + router = _make_router_with_fallback("anthropic/claude-sonnet-4-6", "vertex_ai/claude-sonnet-4-6") src = _make_responses_iterator( chunks=[MagicMock(type="response.created")], error=MidStreamFallbackError( @@ -4644,9 +4547,9 @@ async def test_aresponses_streaming_iterator_writes_litellm_metadata_on_fallback fbk = mock_fallback_utils.call_args.kwargs["kwargs"] assert "litellm_metadata" in fbk, "wrong metadata_variable_name" assert fbk["litellm_metadata"]["model_group"] == "gpt-4" - assert "model_group" not in fbk.get( - "metadata", {} - ), "model_group leaked into 'metadata' instead of 'litellm_metadata'" + assert "model_group" not in fbk.get("metadata", {}), ( + "model_group leaked into 'metadata' instead of 'litellm_metadata'" + ) @pytest.mark.asyncio @@ -5061,9 +4964,7 @@ async def test_aresponses_streaming_iterator_combines_partial_usage(): fallback_response_object = ResponsesAPIResponse( id="resp_test", created_at=0, model="gpt-4", object="response", output=[] ) - fallback_response_object.usage = ResponseAPIUsage( - input_tokens=20, output_tokens=15, total_tokens=35 - ) + fallback_response_object.usage = ResponseAPIUsage(input_tokens=20, output_tokens=15, total_tokens=35) fallback_event = ResponseCompletedEvent( type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, response=fallback_response_object, @@ -5072,9 +4973,7 @@ async def test_aresponses_streaming_iterator_combines_partial_usage(): with ( patch( "litellm.main.stream_chunk_builder", - return_value=SimpleNamespace( - usage=SimpleNamespace(prompt_tokens=10, completion_tokens=4) - ), + return_value=SimpleNamespace(usage=SimpleNamespace(prompt_tokens=10, completion_tokens=4)), ), patch.object( router, @@ -5375,9 +5274,7 @@ def test_pre_call_checks_skips_token_count_without_max_input_tokens(monkeypatch) monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {}) calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 - ) + monkeypatch.setattr(litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000) deployments = [ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, @@ -5405,14 +5302,10 @@ def test_pre_call_checks_counts_once_and_filters_on_max_input_tokens(monkeypatch ], enable_pre_call_checks=True, ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5}) calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 - ) + monkeypatch.setattr(litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000) deployments = [ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, @@ -5439,14 +5332,10 @@ def test_pre_call_checks_uses_precounted_tokens(monkeypatch): ], enable_pre_call_checks=True, ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5}) calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1 - ) + monkeypatch.setattr(litellm, "token_counter", lambda *a, **k: calls.append(1) or 1) deployments = [ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, @@ -5473,9 +5362,7 @@ async def test_async_get_healthy_deployments_counts_tokens_off_the_event_loop(mo ], enable_pre_call_checks=True, ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1_000_000} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1_000_000}) counting_threads = [] monkeypatch.setattr( @@ -5571,14 +5458,10 @@ def test_pre_call_checks_does_not_recount_inline_after_an_off_loop_failure(monke ], enable_pre_call_checks=True, ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5}) calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 - ) + monkeypatch.setattr(litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000) deployments = [ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, @@ -5606,9 +5489,7 @@ async def test_async_get_healthy_deployments_never_recounts_on_the_loop(monkeypa ], enable_pre_call_checks=True, ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5}) counting_threads = [] @@ -5643,9 +5524,7 @@ async def test_acount_pre_call_check_tokens_leaves_the_event_loop_free(monkeypat ], enable_pre_call_checks=True, ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5}) deployments = [ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, @@ -5681,9 +5560,7 @@ async def test_acount_pre_call_check_tokens_skips_without_max_input_tokens(monke monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {}) calls = [] - monkeypatch.setattr( - litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 - ) + monkeypatch.setattr(litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000) count = await router._acount_pre_call_check_tokens( model="m", @@ -5711,9 +5588,7 @@ def test_pre_call_checks_counts_tokens_from_responses_input_string(monkeypatch): ], enable_pre_call_checks=True, ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1}) deployments = [ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, @@ -5738,9 +5613,7 @@ def test_pre_call_checks_counts_tokens_from_responses_input_list(monkeypatch): ], enable_pre_call_checks=True, ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1}) deployments = [ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, @@ -5782,9 +5655,7 @@ def test_pre_call_checks_counts_responses_instructions_tokens(monkeypatch): ) assert with_instructions_tokens > input_only_tokens - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": input_only_tokens} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": input_only_tokens}) with pytest.raises(litellm.ContextWindowExceededError): router._pre_call_checks( model="m", @@ -5853,9 +5724,7 @@ def test_pre_call_checks_counts_tool_definition_tokens(monkeypatch, prompt_kwarg prompt_only_tokens = router._count_pre_call_check_tokens( messages=prompt_kwargs.get("messages"), input=prompt_kwargs.get("input") ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": prompt_only_tokens} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": prompt_only_tokens}) assert len(router._pre_call_checks(model="m", healthy_deployments=deployments, **prompt_kwargs)) == 1 with pytest.raises(litellm.ContextWindowExceededError): @@ -5975,7 +5844,7 @@ def test_count_pre_call_check_tokens_across_api_surfaces(): assert string_input_tokens > 0 assert list_input_tokens > 0 - with pytest.raises(ValueError, match='Either messages or input must be provided to count tokens'): + with pytest.raises(ValueError, match="Either messages or input must be provided to count tokens"): router._count_pre_call_check_tokens(messages=None, input=None) @@ -5990,9 +5859,7 @@ def test_pre_call_checks_no_messages_or_input_does_not_crash(monkeypatch): ], enable_pre_call_checks=True, ) - monkeypatch.setattr( - router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} - ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5}) counted: list[dict] = [] original = router._count_pre_call_check_tokens @@ -6082,9 +5949,7 @@ def test_get_deployment_model_info_base_model_flow(): } # Test Case 1: Base model flow with custom model info that has base_model - with patch.object( - litellm, "model_cost", {"test-custom-model": mock_custom_model_info} - ): + with patch.object(litellm, "model_cost", {"test-custom-model": mock_custom_model_info}): with patch.object(litellm, "get_model_info") as mock_get_model_info: # Configure mock returns mock_get_model_info.side_effect = lambda model: { @@ -6092,15 +5957,11 @@ def test_get_deployment_model_info_base_model_flow(): "test-model": mock_litellm_model_name_info, }.get(model) - result = router.get_deployment_model_info( - model_id="test-custom-model", model_name="test-model" - ) + result = router.get_deployment_model_info(model_id="test-custom-model", model_name="test-model") # Verify that get_model_info was called for both base model and model name assert mock_get_model_info.call_count == 2 - mock_get_model_info.assert_any_call( - model="gpt-3.5-turbo" - ) # base model call + mock_get_model_info.assert_any_call(model="gpt-3.5-turbo") # base model call mock_get_model_info.assert_any_call(model="test-model") # model name call # Verify the result contains merged information @@ -6111,26 +5972,18 @@ def test_get_deployment_model_info_base_model_flow(): # 2. The result of step 1 gets merged into litellm_model_name_info (custom+base override litellm) # Fields from custom model (should override base model values) - assert ( - result["input_cost_per_token"] == 0.001 - ) # From custom model (overrides base 0.0015) - assert ( - result["output_cost_per_token"] == 0.002 - ) # From custom model (same as base) + assert result["input_cost_per_token"] == 0.001 # From custom model (overrides base 0.0015) + assert result["output_cost_per_token"] == 0.002 # From custom model (same as base) assert result["custom_field"] == "custom_value" # From custom model # Fields from base model that weren't overridden by custom assert result["max_tokens"] == 4096 # From base model assert result["litellm_provider"] == "openai" # From base model - assert ( - result["mode"] == "chat" - ) # From base model (overrides litellm "completion") + assert result["mode"] == "chat" # From base model (overrides litellm "completion") # The key field comes from base model since both base and litellm have it # and base model info overrides litellm model name info in final merge - assert ( - result["key"] == "gpt-3.5-turbo" - ) # From base model (overrides litellm key) + assert result["key"] == "gpt-3.5-turbo" # From base model (overrides litellm key) # Test Case 2: Custom model info without base_model mock_custom_model_info_no_base = { @@ -6149,9 +6002,7 @@ def test_get_deployment_model_info_base_model_flow(): "test-model": mock_litellm_model_name_info, }.get(model) - result = router.get_deployment_model_info( - model_id="test-custom-model-no-base", model_name="test-model" - ) + result = router.get_deployment_model_info(model_id="test-custom-model-no-base", model_name="test-model") # Should only call get_model_info once for model name (no base model) assert mock_get_model_info.call_count == 1 @@ -6171,9 +6022,7 @@ def test_get_deployment_model_info_base_model_flow(): "test-model": mock_litellm_model_name_info, }.get(model) - result = router.get_deployment_model_info( - model_id="non-existent-model", model_name="test-model" - ) + result = router.get_deployment_model_info(model_id="non-existent-model", model_name="test-model") # Should only call get_model_info once for model name assert mock_get_model_info.call_count == 1 @@ -6206,9 +6055,7 @@ def test_get_deployment_model_info_base_model_flow(): mock_get_model_info.side_effect = mock_get_model_info_side_effect - result = router.get_deployment_model_info( - model_id="test-custom-model-invalid", model_name="test-model" - ) + result = router.get_deployment_model_info(model_id="test-custom-model-invalid", model_name="test-model") # Should handle exception gracefully and still return merged result assert result is not None @@ -6217,12 +6064,8 @@ def test_get_deployment_model_info_base_model_flow(): # Test Case 5: Both model_cost.get() and get_model_info() return None with patch.object(litellm, "model_cost", {}): - with patch.object( - litellm, "get_model_info", side_effect=Exception("Not found") - ): - result = router.get_deployment_model_info( - model_id="non-existent", model_name="non-existent" - ) + with patch.object(litellm, "get_model_info", side_effect=Exception("Not found")): + result = router.get_deployment_model_info(model_id="non-existent", model_name="non-existent") # Should return None when no model info is found assert result is None @@ -6245,9 +6088,7 @@ def test_get_deployment_model_info_base_model_flow(): # Model NOT in built-in cost map — raise exception mock_get_model_info.side_effect = Exception("Model not in cost map") - result = router.get_deployment_model_info( - model_id="custom-model-id", model_name="unknown-model" - ) + result = router.get_deployment_model_info(model_id="custom-model-id", model_name="unknown-model") # Should return custom_model_info even when litellm_model_name_model_info is None assert result is not None @@ -6283,15 +6124,11 @@ def test_get_deployment_model_info_base_model_flow(): mock_get_model_info.side_effect = get_info_side_effect - result = router.get_deployment_model_info( - model_id="custom-with-base", model_name="unknown-model" - ) + result = router.get_deployment_model_info(model_id="custom-with-base", model_name="unknown-model") # Should return custom_model_info merged with base model info assert result is not None - assert ( - result["input_cost_per_token"] == 0.01 - ) # From custom (overrides base) + assert result["input_cost_per_token"] == 0.01 # From custom (overrides base) assert result["max_tokens"] == 8192 # From base model assert result["litellm_provider"] == "openai" # From base model @@ -6338,18 +6175,14 @@ def test_get_deployment_model_info_base_model_merge_priority(): "litellm_only_field": "litellm_value", } - with patch.object( - litellm, "model_cost", {"custom-model-id": mock_custom_model_info} - ): + with patch.object(litellm, "model_cost", {"custom-model-id": mock_custom_model_info}): with patch.object(litellm, "get_model_info") as mock_get_model_info: mock_get_model_info.side_effect = lambda model: { "gpt-4": mock_base_model_info, "test-model": mock_litellm_model_name_info, }.get(model) - result = router.get_deployment_model_info( - model_id="custom-model-id", model_name="test-model" - ) + result = router.get_deployment_model_info(model_id="custom-model-id", model_name="test-model") assert result is not None @@ -6359,29 +6192,17 @@ def test_get_deployment_model_info_base_model_merge_priority(): # 3. Result from steps 1-2 overrides litellm_model_name_info # Fields that should come from custom model info (highest priority) - assert ( - result["input_cost_per_token"] == 0.01 - ) # From custom model (overrides base 0.03) - assert ( - result["max_tokens"] == 8000 - ) # From custom model (overrides base 4096) + assert result["input_cost_per_token"] == 0.01 # From custom model (overrides base 0.03) + assert result["max_tokens"] == 8000 # From custom model (overrides base 4096) assert result["custom_only_field"] == "custom_value" # From custom model # Fields that should come from base model (not overridden by custom) - assert ( - result["output_cost_per_token"] == 0.06 - ) # From base model (not in custom) - assert ( - result["litellm_provider"] == "openai" - ) # From base model (not in custom) - assert ( - result["base_only_field"] == "base_value" - ) # From base model (not in custom) + assert result["output_cost_per_token"] == 0.06 # From base model (not in custom) + assert result["litellm_provider"] == "openai" # From base model (not in custom) + assert result["base_only_field"] == "base_value" # From base model (not in custom) # Fields that should come from litellm model name info (not overridden by custom+base) - assert ( - result["mode"] == "completion" - ) # From litellm model name info (not in custom or base) + assert result["mode"] == "completion" # From litellm model name info (not in custom or base) assert ( result["litellm_only_field"] == "litellm_value" ) # From litellm model name info (not in custom or base) @@ -6398,7 +6219,11 @@ def test_get_deployment_model_info_base_model_merge_priority(): [ ( "gpt", - {"model": "azure_ai/gpt-5.4-mini", "api_base": "https://my-resource.services.ai.azure.com", "api_key": "key"}, + { + "model": "azure_ai/gpt-5.4-mini", + "api_base": "https://my-resource.services.ai.azure.com", + "api_key": "key", + }, "gpt/openai/deployments/gpt-5.4-mini/chat/completions", "gpt-5.4-mini/openai/deployments/gpt-5.4-mini/chat/completions", ), @@ -6453,10 +6278,9 @@ def test_add_deployment_model_to_endpoint_for_llm_passthrough_route(): model="special-bedrock-model", model_name="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", ) - assert ( - result["endpoint"] - == "/model/us.anthropic.claude-haiku-4-5-20251001-v1:0/invoke" - ), f"Expected '/model/us.anthropic.claude-haiku-4-5-20251001-v1:0/invoke', got '{result['endpoint']}'" + assert result["endpoint"] == "/model/us.anthropic.claude-haiku-4-5-20251001-v1:0/invoke", ( + f"Expected '/model/us.anthropic.claude-haiku-4-5-20251001-v1:0/invoke', got '{result['endpoint']}'" + ) # Test Case 2: Bedrock invoke-with-response-stream endpoint kwargs = { @@ -6468,10 +6292,9 @@ def test_add_deployment_model_to_endpoint_for_llm_passthrough_route(): model="special-bedrock-model", model_name="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", ) - assert ( - result["endpoint"] - == "/model/us.anthropic.claude-haiku-4-5-20251001-v1:0/invoke-with-response-stream" - ), f"Expected streaming endpoint with stripped prefix, got '{result['endpoint']}'" + assert result["endpoint"] == "/model/us.anthropic.claude-haiku-4-5-20251001-v1:0/invoke-with-response-stream", ( + f"Expected streaming endpoint with stripped prefix, got '{result['endpoint']}'" + ) # Test Case 3: Bedrock converse endpoint kwargs = { @@ -6483,9 +6306,9 @@ def test_add_deployment_model_to_endpoint_for_llm_passthrough_route(): model="bedrock-model", model_name="bedrock/us.meta.llama3-8b-instruct-v1:0", ) - assert ( - result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/converse" - ), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/converse', got '{result['endpoint']}'" + assert result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/converse", ( + f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/converse', got '{result['endpoint']}'" + ) # Test Case 4: Bedrock provider prefix auto-detected from model_name kwargs = { @@ -6496,9 +6319,9 @@ def test_add_deployment_model_to_endpoint_for_llm_passthrough_route(): model="router-model", model_name="bedrock/us.meta.llama3-8b-instruct-v1:0", ) - assert ( - result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/invoke" - ), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/invoke', got '{result['endpoint']}'" + assert result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/invoke", ( + f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/invoke', got '{result['endpoint']}'" + ) def test_update_kwargs_with_deployment_uses_pass_through_request_timeout(): @@ -6635,14 +6458,10 @@ async def test_router_acompletion_with_unknown_model_and_default_fallback(): # Initialize the router with a default fallback router = litellm.Router(model_list=model_list, default_fallbacks=["gpt-4o"]) - messages = [ - {"role": "user", "content": "This call should succeed by falling back."} - ] + messages = [{"role": "user", "content": "This call should succeed by falling back."}] # Call completion with a model name that is NOT in the model_list - response = await router.acompletion( - model="completely-unknown-model", messages=messages - ) + response = await router.acompletion(model="completely-unknown-model", messages=messages) # Check that the call did not fail and we received a valid response object. assert response is not None @@ -6821,15 +6640,10 @@ def test_get_deployment_credentials_with_provider_aws_bedrock_runtime_endpoint() ], ) - credentials = router.get_deployment_credentials_with_provider( - model_id="bedrock-claude-model" - ) + credentials = router.get_deployment_credentials_with_provider(model_id="bedrock-claude-model") assert credentials is not None - assert ( - credentials["aws_bedrock_runtime_endpoint"] - == "https://bedrock-runtime.us-east-1.amazonaws.com" - ) + assert credentials["aws_bedrock_runtime_endpoint"] == "https://bedrock-runtime.us-east-1.amazonaws.com" assert credentials["aws_access_key_id"] == "test-access-key" assert credentials["aws_secret_access_key"] == "test-secret-key" assert credentials["aws_region_name"] == "us-east-1" @@ -6856,9 +6670,7 @@ def test_get_deployment_credentials_with_provider_includes_bucket_name(): ], ) - credentials = router.get_deployment_credentials_with_provider( - model_id="vertex-gemini" - ) + credentials = router.get_deployment_credentials_with_provider(model_id="vertex-gemini") assert credentials is not None assert credentials["gcs_bucket_name"] == "my-batch-bucket" @@ -6943,9 +6755,7 @@ def test_get_deployment_credentials_with_provider_resolves_credential_name(): ], ) - credentials = router.get_deployment_credentials_with_provider( - model_id="azure-gpt-4" - ) + credentials = router.get_deployment_credentials_with_provider(model_id="azure-gpt-4") assert credentials is not None assert credentials["api_key"] == "resolved-api-key" @@ -6982,9 +6792,7 @@ def test_get_deployment_credentials_with_provider_bedrock_batch_fields(): ], ) - credentials = router.get_deployment_credentials_with_provider( - model_id="bedrock-batch-model" - ) + credentials = router.get_deployment_credentials_with_provider(model_id="bedrock-batch-model") assert credentials is not None assert credentials["custom_llm_provider"] == "bedrock" @@ -7028,9 +6836,7 @@ def test_get_deployment_credentials_with_provider_preserves_aws_auth_params(): ], ) - credentials = router.get_deployment_credentials_with_provider( - model_id="bedrock-batch-model" - ) + credentials = router.get_deployment_credentials_with_provider(model_id="bedrock-batch-model") assert credentials is not None for key, value in aws_auth_params.items(): @@ -7100,15 +6906,11 @@ def test_get_deployment_credentials_with_provider_team_wildcard_priority(): ], ) - team_credentials = router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) + team_credentials = router.get_deployment_credentials_with_provider(model_id="openai/gpt-5.2", team_id="team-1") assert team_credentials is not None assert team_credentials["api_key"] == "team-key" - global_credentials = router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2" - ) + global_credentials = router.get_deployment_credentials_with_provider(model_id="openai/gpt-5.2") assert global_credentials is not None assert global_credentials["api_key"] == "global-key" @@ -7149,15 +6951,11 @@ def test_get_deployment_credentials_with_provider_skips_other_team_deployment(): assert other_team_credentials is not None assert other_team_credentials["vertex_project"] == "shared-project" - unscoped_credentials = router.get_deployment_credentials_with_provider( - model_id="gemini-2.5-pro" - ) + unscoped_credentials = router.get_deployment_credentials_with_provider(model_id="gemini-2.5-pro") assert unscoped_credentials is not None assert unscoped_credentials["vertex_project"] == "shared-project" - owner_credentials = router.get_deployment_credentials_with_provider( - model_id="gemini-2.5-pro", team_id="team-b" - ) + owner_credentials = router.get_deployment_credentials_with_provider(model_id="gemini-2.5-pro", team_id="team-b") assert owner_credentials is not None assert owner_credentials["vertex_project"] == "team-b-project" @@ -7184,16 +6982,8 @@ def test_get_deployment_credentials_with_provider_no_fallback_to_other_team_only ], ) - assert ( - router.get_deployment_credentials_with_provider( - model_id="gemini-2.5-pro", team_id="team-a" - ) - is None - ) - assert ( - router.get_deployment_credentials_with_provider(model_id="gemini-2.5-pro") - is None - ) + assert router.get_deployment_credentials_with_provider(model_id="gemini-2.5-pro", team_id="team-a") is None + assert router.get_deployment_credentials_with_provider(model_id="gemini-2.5-pro") is None def test_deployment_usable_by_team_helpers(): @@ -7233,9 +7023,7 @@ def test_deployment_usable_by_team_helpers(): assert router._deployment_usable_by_team(shared, "team-a") is True assert router._deployment_usable_by_team(shared, None) is True - picked = router._get_model_group_deployment_usable_by_team( - model_group_name="gemini-2.5-pro", team_id="team-a" - ) + picked = router._get_model_group_deployment_usable_by_team(model_group_name="gemini-2.5-pro", team_id="team-a") assert picked is not None assert picked.litellm_params.vertex_project == "shared-project" @@ -7245,12 +7033,15 @@ def test_deployment_usable_by_team_helpers(): assert owner_picked is not None assert owner_picked.litellm_params.vertex_project == "team-b-project" - assert ( - router._get_model_group_deployment_usable_by_team( - model_group_name="unknown-model", team_id="team-a" - ) - is None - ) + assert router._get_model_group_deployment_usable_by_team(model_group_name="unknown-model", team_id="team-a") is None + + +def test_deployment_usable_by_team_uses_dynamic_model_info_get(): + class ModelInfo: + def get(self, key: str) -> str | None: + return "team-b" if key == "team_id" else None + + assert litellm.Router.deployment_usable_by_team({"model_info": ModelInfo()}, "team-a") is False def test_get_deployment_credentials_with_provider_skips_other_team_wildcard(): @@ -7282,9 +7073,7 @@ def test_get_deployment_credentials_with_provider_skips_other_team_wildcard(): assert other_team_credentials is not None assert other_team_credentials["api_key"] == "global-key" - owner_credentials = router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-b" - ) + owner_credentials = router.get_deployment_credentials_with_provider(model_id="openai/gpt-5.2", team_id="team-b") assert owner_credentials is not None assert owner_credentials["api_key"] == "team-b-key" @@ -7296,21 +7085,11 @@ def test_team_wildcard_credentials_not_usable_after_delete_deployment(): """ router = litellm.Router(model_list=[_team_wildcard_model(api_key="old-key")]) - assert ( - router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) - is not None - ) + assert router.get_deployment_credentials_with_provider(model_id="openai/gpt-5.2", team_id="team-1") is not None router.delete_deployment(id="team-wildcard-id") - assert ( - router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) - is None - ) + assert router.get_deployment_credentials_with_provider(model_id="openai/gpt-5.2", team_id="team-1") is None def test_global_wildcard_pattern_router_evicts_stale_entry_on_upsert_and_delete(): @@ -7383,22 +7162,13 @@ def test_team_wildcard_credentials_refreshed_on_upsert_and_set_model_list(): router = litellm.Router(model_list=[_team_wildcard_model(api_key="old-key")]) - router.upsert_deployment( - deployment=Deployment(**_team_wildcard_model(api_key="new-key")) - ) - credentials = router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) + router.upsert_deployment(deployment=Deployment(**_team_wildcard_model(api_key="new-key"))) + credentials = router.get_deployment_credentials_with_provider(model_id="openai/gpt-5.2", team_id="team-1") assert credentials is not None assert credentials["api_key"] == "new-key" router.set_model_list(model_list=[]) - assert ( - router.get_deployment_credentials_with_provider( - model_id="openai/gpt-5.2", team_id="team-1" - ) - is None - ) + assert router.get_deployment_credentials_with_provider(model_id="openai/gpt-5.2", team_id="team-1") is None def test_get_available_guardrail_single_deployment(): @@ -7587,9 +7357,7 @@ async def test_anthropic_messages_call_type_is_cached(): startTime=1234567890.0, endTime=1234567891.0, completionStartTime=1234567890.5, - model_map_information=StandardLoggingModelInformation( - model_map_key="gpt-3.5-turbo", model_map_value=None - ), + model_map_information=StandardLoggingModelInformation(model_map_key="gpt-3.5-turbo", model_map_value=None), model="gpt-3.5-turbo", model_id="model-123", model_group="openai-gpt", @@ -7668,12 +7436,8 @@ async def test_anthropic_messages_call_type_is_cached(): ) # This assertion will FAIL if anthropic_messages is filtered out - assert ( - cached_result is not None - ), "Model ID should be cached for anthropic_messages call type" - assert ( - cached_result["model_id"] == test_model_id - ), f"Expected {test_model_id}, got {cached_result['model_id']}" + assert cached_result is not None, "Model ID should be cached for anthropic_messages call type" + assert cached_result["model_id"] == test_model_id, f"Expected {test_model_id}, got {cached_result['model_id']}" def test_update_kwargs_with_deployment_propagates_model_tags(): @@ -7698,9 +7462,7 @@ def test_update_kwargs_with_deployment_propagates_model_tags(): ) kwargs: dict = {"metadata": {}} - deployment = router.get_deployment_by_model_group_name( - model_group_name="gpt-4o-mini" - ) + deployment = router.get_deployment_by_model_group_name(model_group_name="gpt-4o-mini") router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) # Deployment tags should be propagated to kwargs metadata @@ -7729,9 +7491,7 @@ def test_update_kwargs_with_deployment_merges_tags_without_duplicates(): # Simulate request that already has tags (from request body or key/team level) kwargs: dict = {"metadata": {"tags": ["user-tag", "shared-tag"]}} - deployment = router.get_deployment_by_model_group_name( - model_group_name="gpt-4o-mini" - ) + deployment = router.get_deployment_by_model_group_name(model_group_name="gpt-4o-mini") router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) # Both sources should be merged, no duplicates @@ -7758,9 +7518,7 @@ def test_update_kwargs_with_deployment_no_tags(): ) kwargs: dict = {"metadata": {}} - deployment = router.get_deployment_by_model_group_name( - model_group_name="gpt-4o-mini" - ) + deployment = router.get_deployment_by_model_group_name(model_group_name="gpt-4o-mini") router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) # No tags key should be added if deployment has no tags @@ -7842,9 +7600,7 @@ def test_update_kwargs_with_deployment_merges_tools(): }, ], } - deployment = router.get_deployment_by_model_group_name( - model_group_name="o3-deep-research" - ) + deployment = router.get_deployment_by_model_group_name(model_group_name="o3-deep-research") router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) # Tools should be merged: deployment first, then request @@ -7875,9 +7631,7 @@ def test_update_kwargs_with_deployment_merge_tools_deployment_only(): ) kwargs: dict = {"metadata": {}} - deployment = router.get_deployment_by_model_group_name( - model_group_name="o3-deep-research" - ) + deployment = router.get_deployment_by_model_group_name(model_group_name="o3-deep-research") router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) assert kwargs["tools"] == [{"type": "web_search"}] @@ -7906,9 +7660,7 @@ def test_update_kwargs_with_deployment_merge_tools_request_overrides_tool_choice "metadata": {}, "tool_choice": "none", } - deployment = router.get_deployment_by_model_group_name( - model_group_name="o3-deep-research" - ) + deployment = router.get_deployment_by_model_group_name(model_group_name="o3-deep-research") router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) # Request tool_choice should be preserved (merged tools still applied) @@ -8010,12 +7762,8 @@ def test_update_kwargs_with_deployment_model_info_in_litellm_metadata(): ) kwargs: dict = {} - deployment = router.get_deployment_by_model_group_name( - model_group_name="claude-sonnet-4" - ) - router._update_kwargs_with_deployment( - deployment=deployment, kwargs=kwargs, function_name="generic_api_call" - ) + deployment = router.get_deployment_by_model_group_name(model_group_name="claude-sonnet-4") + router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name="generic_api_call") assert "litellm_metadata" in kwargs model_info = kwargs["litellm_metadata"]["model_info"] @@ -8047,12 +7795,8 @@ def test_update_kwargs_with_deployment_model_info_in_metadata(): ) kwargs: dict = {} - deployment = router.get_deployment_by_model_group_name( - model_group_name="claude-sonnet-4" - ) - router._update_kwargs_with_deployment( - deployment=deployment, kwargs=kwargs, function_name=None - ) + deployment = router.get_deployment_by_model_group_name(model_group_name="claude-sonnet-4") + router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name=None) assert "metadata" in kwargs model_info = kwargs["metadata"]["model_info"] @@ -8165,6 +7909,7 @@ async def test_acompletion_streaming_iterator_does_not_log_success_on_terminal_f initial_kwargs=dict(initial_kwargs), ) collected = [] + async def _drain(): async for chunk in result: collected.append(chunk) @@ -8191,6 +7936,7 @@ async def test_acompletion_streaming_iterator_does_not_log_success_on_terminal_f initial_kwargs=dict(initial_kwargs), ) collected = [] + async def _drain(): async for chunk in result: collected.append(chunk) @@ -8407,23 +8153,17 @@ def test_multiregion_team_deployments_unique_model_names(): assert len(deployments) == 0 # With team_id: O(n) scan finds BOTH regional deployments - deployments = router._get_all_deployments( - model_name="claude-sonnet", team_id="metis-team" - ) + deployments = router._get_all_deployments(model_name="claude-sonnet", team_id="metis-team") assert len(deployments) == 2 deployment_names = {d["model_name"] for d in deployments} assert deployment_names == {"metis-claude-us-east-1", "metis-claude-us-west-2"} # Each deployment has a unique ID (critical for cooldown/retry to work) deployment_ids = {d["model_info"]["id"] for d in deployments} - assert ( - len(deployment_ids) == 2 - ), "Each deployment must have a unique ID for cooldown tracking" + assert len(deployment_ids) == 2, "Each deployment must have a unique ID for cooldown tracking" # Wrong team: returns nothing - deployments = router._get_all_deployments( - model_name="claude-sonnet", team_id="other-team" - ) + deployments = router._get_all_deployments(model_name="claude-sonnet", team_id="other-team") assert len(deployments) == 0 @@ -8468,12 +8208,8 @@ async def test_multiregion_team_failover_between_regions(): ) # Verify the router finds both deployments for the team - deployments = router._get_all_deployments( - model_name="claude-sonnet", team_id="metis-team" - ) - assert ( - len(deployments) == 2 - ), "Router must find both regional deployments by team_public_model_name" + deployments = router._get_all_deployments(model_name="claude-sonnet", team_id="metis-team") + assert len(deployments) == 2, "Router must find both regional deployments by team_public_model_name" # Make a normal request — should succeed from one of the regions response = await router.acompletion( @@ -8598,9 +8334,7 @@ def test_explicit_model_access_does_not_force_access_group_filtering(): }, ) - deployment_groups = [ - d.get("model_info", {}).get("access_groups") for d in deployments - ] + deployment_groups = [d.get("model_info", {}).get("access_groups") for d in deployments] assert ["AG1"] in deployment_groups assert ["AG2"] in deployment_groups @@ -8645,9 +8379,7 @@ def test_access_group_filter_empty_does_not_bypass_via_litellm_model_fallback( orig_groups = router.get_model_access_groups - def fake_get_model_access_groups( - model_name=None, model_access_group=None, team_id=None - ): + def fake_get_model_access_groups(model_name=None, model_access_group=None, team_id=None): if model_name == "gpt-5" and model_access_group is None: return {"AG1": ["gpt-5"], "AG2": ["gpt-5"]} return orig_groups( @@ -8722,9 +8454,7 @@ def test_access_group_block_does_not_silently_use_default_fallback_model( orig_groups = router.get_model_access_groups - def fake_get_model_access_groups( - model_name=None, model_access_group=None, team_id=None - ): + def fake_get_model_access_groups(model_name=None, model_access_group=None, team_id=None): if model_name == "gpt-5" and model_access_group is None: return {"AG1": ["gpt-5"], "AG2": ["gpt-5"]} return orig_groups( @@ -8791,9 +8521,7 @@ def test_access_group_block_via_litellm_model_branch_does_not_use_default_fallba orig_groups = router.get_model_access_groups - def fake_get_model_access_groups( - model_name=None, model_access_group=None, team_id=None - ): + def fake_get_model_access_groups(model_name=None, model_access_group=None, team_id=None): if model_name == "gpt-5" and model_access_group is None: return {"AG1": ["gpt-5"], "AG2": ["gpt-5"]} return orig_groups( @@ -8848,9 +8576,7 @@ def test_try_early_resolve_deployments_for_model_not_in_names(): ) assert ( - router_in_names._try_early_resolve_deployments_for_model_not_in_names( - model="gpt-5", request_team_id=None - ) + router_in_names._try_early_resolve_deployments_for_model_not_in_names(model="gpt-5", request_team_id=None) is None ) assert ( @@ -8872,10 +8598,8 @@ def test_try_early_resolve_deployments_for_model_not_in_names(): ] ) - pattern_result = ( - pattern_router._try_early_resolve_deployments_for_model_not_in_names( - model="openai/gpt-4o-mini", request_team_id=None - ) + pattern_result = pattern_router._try_early_resolve_deployments_for_model_not_in_names( + model="openai/gpt-4o-mini", request_team_id=None ) assert pattern_result is not None resolved_model, pattern_deployments = pattern_result @@ -8901,10 +8625,8 @@ def test_try_early_resolve_deployments_for_model_not_in_names(): }, } - default_result = ( - default_router._try_early_resolve_deployments_for_model_not_in_names( - model="brand-new-model", request_team_id=None - ) + default_result = default_router._try_early_resolve_deployments_for_model_not_in_names( + model="brand-new-model", request_team_id=None ) assert default_result is not None resolved_model, default_deployment = default_result @@ -8912,10 +8634,7 @@ def test_try_early_resolve_deployments_for_model_not_in_names(): assert isinstance(default_deployment, dict) assert default_deployment["litellm_params"]["model"] == "brand-new-model" # The original default_deployment must not be mutated. - assert ( - default_router.default_deployment["litellm_params"]["model"] - == "openai/will-be-overridden" - ) + assert default_router.default_deployment["litellm_params"]["model"] == "openai/will-be-overridden" def _router_with_two_deployments(blocked_flags): @@ -8963,10 +8682,7 @@ def _seed_unhealthy_states(router, unhealthy_ids, timestamp=None): ts = timestamp if timestamp is not None else time.time() router.health_state_cache.set_deployment_health_states( - { - uid: {"is_healthy": False, "timestamp": ts, "reason": "test_unhealthy"} - for uid in unhealthy_ids - } + {uid: {"is_healthy": False, "timestamp": ts, "reason": "test_unhealthy"} for uid in unhealthy_ids} ) @@ -9033,7 +8749,6 @@ async def test_health_probe_preserves_normal_caller_policy( assert await router.cooldown_cache.async_get_active_cooldowns(["dep-0", "dep-1"], parent_otel_span=None) == [] - @pytest.mark.asyncio async def test_async_get_fully_unhealthy_model_names_keeps_name_when_partial(): router = _router_with_two_deployments([False, False]) @@ -9094,9 +8809,7 @@ async def test_async_get_fully_unhealthy_model_names_noop_with_allowed_fails_pol @pytest.mark.asyncio async def test_async_get_healthy_deployments_skips_blocked_deployment(): router = _router_with_two_deployments([True, False]) - healthy, all_dep = await router._async_get_healthy_deployments( - model="gpt-4o", parent_otel_span=None - ) + healthy, all_dep = await router._async_get_healthy_deployments(model="gpt-4o", parent_otel_span=None) healthy_ids = [d["model_info"]["id"] for d in healthy] assert "dep-0" not in healthy_ids assert "dep-1" in healthy_ids @@ -9105,9 +8818,7 @@ async def test_async_get_healthy_deployments_skips_blocked_deployment(): def test_get_healthy_deployments_sync_skips_blocked_deployment(): router = _router_with_two_deployments([False, True]) - healthy, all_dep = router._get_healthy_deployments( - model="gpt-4o", parent_otel_span=None - ) + healthy, all_dep = router._get_healthy_deployments(model="gpt-4o", parent_otel_span=None) healthy_ids = [d["model_info"]["id"] for d in healthy] assert "dep-0" in healthy_ids assert "dep-1" not in healthy_ids @@ -9124,9 +8835,7 @@ def test_filter_blocked_deployments_drops_blocked_keeps_unblocked(): @pytest.mark.asyncio async def test_public_async_get_healthy_deployments_skips_blocked_on_primary_path(): router = _router_with_two_deployments([True, False]) - deployments = await router.async_get_healthy_deployments( - model="gpt-4o", request_kwargs={} - ) + deployments = await router.async_get_healthy_deployments(model="gpt-4o", request_kwargs={}) assert isinstance(deployments, list) ids = [d["model_info"]["id"] for d in deployments] assert "dep-0" not in ids @@ -9214,9 +8923,7 @@ def _router_with_two_pass_through_deployments(blocked_flags): def test_get_available_deployment_for_pass_through_skips_blocked(): router = _router_with_two_pass_through_deployments([True, False]) - deployment = router.get_available_deployment_for_pass_through( - model="gpt-4o", request_kwargs={} - ) + deployment = router.get_available_deployment_for_pass_through(model="gpt-4o", request_kwargs={}) assert deployment["model_info"]["id"] == "pt-1" @@ -9225,9 +8932,7 @@ def test_get_available_deployment_for_pass_through_raises_when_dict_blocked(): router = _router_with_two_pass_through_deployments([True, True]) with pytest.raises(litellm.ServiceUnavailableError): - router.get_available_deployment_for_pass_through( - model="pt-0", request_kwargs={} - ) + router.get_available_deployment_for_pass_through(model="pt-0", request_kwargs={}) def test_get_available_deployment_for_pass_through_names_cooldown_despite_healthy_non_pass_through(): @@ -9269,9 +8974,7 @@ def test_initialize_deployment_for_pass_through_keeps_bedrock_iam_deployment(): } ] ) - assert [m["model_info"]["id"] for m in router.get_model_list()] == [ - "bedrock-iam-pt" - ] + assert [m["model_info"]["id"] for m in router.get_model_list()] == ["bedrock-iam-pt"] def test_pass_through_deployment_api_key_resolves_via_get_credentials(): @@ -9282,12 +8985,7 @@ def test_pass_through_deployment_api_key_resolves_via_get_credentials(): router = _router_with_two_pass_through_deployments([False, False]) passthrough_router = PassthroughEndpointRouter(llm_router_getter=lambda: router) assert len(router.get_model_list()) == 2 - assert ( - passthrough_router.get_credentials( - custom_llm_provider="openai", region_name=None - ) - == "sk-fake-for-tests" - ) + assert passthrough_router.get_credentials(custom_llm_provider="openai", region_name=None) == "sk-fake-for-tests" def test_get_deployment_credentials_returns_none_for_blocked_deployment(): @@ -9304,7 +9002,7 @@ def test_get_deployment_credentials_with_provider_returns_none_for_blocked_deplo def test_is_deployment_blocked_static_helper_reflects_blocked_flag(): """ - Exercises Router._is_deployment_blocked so router_code_coverage.py (AST call graph) + Exercises Router.is_deployment_blocked so router_code_coverage.py (AST call graph) marks the helper as covered by router-named tests. """ import types @@ -9315,26 +9013,30 @@ def test_is_deployment_blocked_static_helper_reflects_blocked_flag(): blocked_dep = router.get_deployment("dep-0") unblocked_dep = router.get_deployment("dep-1") assert blocked_dep is not None and unblocked_dep is not None - assert litellm.Router._is_deployment_blocked(blocked_dep) is True - assert litellm.Router._is_deployment_blocked(unblocked_dep) is False + assert litellm.Router.is_deployment_blocked(blocked_dep) is True + assert litellm.Router.is_deployment_blocked(unblocked_dep) is False # No model_info on deployment object → treated as not blocked - assert litellm.Router._is_deployment_blocked(object()) is False + assert litellm.Router.is_deployment_blocked(object()) is False missing_blocked = types.SimpleNamespace() + assert litellm.Router.is_deployment_blocked(types.SimpleNamespace(model_info=missing_blocked)) is False assert ( - litellm.Router._is_deployment_blocked( - types.SimpleNamespace(model_info=missing_blocked) - ) - is False - ) - assert ( - litellm.Router._is_deployment_blocked( - types.SimpleNamespace(model_info=types.SimpleNamespace(blocked=True)) - ) + litellm.Router.is_deployment_blocked(types.SimpleNamespace(model_info=types.SimpleNamespace(blocked=True))) is True ) +def test_routing_groups_legacy_property_forwards_to_public_attribute(): + router = _router_with_two_deployments([False, False]) + + assert router._routing_groups is router.routing_groups + + replacement = dict(router.routing_groups) + router._routing_groups = replacement + + assert router.routing_groups is replacement + + class TestRouterRequestTimeoutPropagation: """litellm_settings.request_timeout must act as an independent per-attempt timeout. @@ -9370,9 +9072,7 @@ class TestRouterRequestTimeoutPropagation: litellm.request_timeout = original_value litellm.request_timeout_explicitly_set = original_flag - def test_request_timeout_stored_independently_when_both_set( - self, explicit_request_timeout - ): + def test_request_timeout_stored_independently_when_both_set(self, explicit_request_timeout): router = self._make_router(timeout=330) assert router.timeout == 330 assert router.request_timeout == 300 @@ -9390,22 +9090,16 @@ class TestRouterRequestTimeoutPropagation: litellm.request_timeout = original_value litellm.request_timeout_explicitly_set = original_flag - def test_non_stream_prefers_request_timeout_over_router_timeout( - self, explicit_request_timeout - ): + def test_non_stream_prefers_request_timeout_over_router_timeout(self, explicit_request_timeout): router = self._make_router(timeout=330) assert router._get_non_stream_timeout(kwargs={}, data={}) == 300 - def test_stream_prefers_request_timeout_over_router_timeout( - self, explicit_request_timeout - ): + def test_stream_prefers_request_timeout_over_router_timeout(self, explicit_request_timeout): router = self._make_router(timeout=330) # stream=True resolves through _get_stream_timeout; request_timeout must win. assert router._get_timeout(kwargs={"stream": True}, data={}) == 300 - def test_explicit_stream_timeout_still_wins_over_request_timeout( - self, explicit_request_timeout - ): + def test_explicit_stream_timeout_still_wins_over_request_timeout(self, explicit_request_timeout): router = self._make_router(timeout=330, stream_timeout=45) assert router._get_stream_timeout(kwargs={}, data={}) == 45 @@ -9421,22 +9115,13 @@ class TestRouterRequestTimeoutPropagation: litellm.request_timeout = original_value litellm.request_timeout_explicitly_set = original_flag - def test_per_deployment_timeout_overrides_request_timeout( - self, explicit_request_timeout - ): + def test_per_deployment_timeout_overrides_request_timeout(self, explicit_request_timeout): router = self._make_router(timeout=330) assert router._get_non_stream_timeout(kwargs={}, data={"timeout": 120}) == 120 - def test_per_request_timeout_overrides_request_timeout( - self, explicit_request_timeout - ): + def test_per_request_timeout_overrides_request_timeout(self, explicit_request_timeout): router = self._make_router(timeout=330) - assert ( - router._get_non_stream_timeout( - kwargs={"timeout": 60}, data={"timeout": 120} - ) - == 60 - ) + assert router._get_non_stream_timeout(kwargs={"timeout": 60}, data={"timeout": 120}) == 60 def test_passthrough_prefers_request_timeout_over_router_timeout(self, explicit_request_timeout): router = self._make_router(timeout=330) @@ -9755,9 +9440,7 @@ class TestAdvisorSubCallCooldown: ) def _cooled_down_ids(self, router): - active = router.cooldown_cache.get_active_cooldowns( - model_ids=["dep-1"], parent_otel_span=None - ) + active = router.cooldown_cache.get_active_cooldowns(model_ids=["dep-1"], parent_otel_span=None) return [entry[0] for entry in active] @pytest.mark.asyncio @@ -9766,12 +9449,7 @@ class TestAdvisorSubCallCooldown: router = self._router() now = datetime.now() - assert ( - router.deployment_callback_on_failure( - self._kwargs(self._auth_error()), None, now, now - ) - is True - ) + assert router.deployment_callback_on_failure(self._kwargs(self._auth_error()), None, now, now) is True assert "dep-1" in self._cooled_down_ids(router) def test_advisor_orchestration_failure_does_not_cool_down_deployment(self): @@ -9786,12 +9464,7 @@ class TestAdvisorSubCallCooldown: mark_advisor_orchestration_failure(exception) now = datetime.now() - assert ( - router.deployment_callback_on_failure( - self._kwargs(exception), None, now, now - ) - is False - ) + assert router.deployment_callback_on_failure(self._kwargs(exception), None, now, now) is False assert "dep-1" not in self._cooled_down_ids(router) @@ -10070,13 +9743,13 @@ def test_stream_chunks_have_generated_content_detects_text_and_non_text(): audio_chunk = _chunk(audio_delta) assert _stream_chunks_have_generated_content([audio_chunk]) is True - images_delta = Delta(images=[{"image_url": {"url": "https://example.com/img.png"}, "index": 0, "type": "image_url"}]) + images_delta = Delta( + images=[{"image_url": {"url": "https://example.com/img.png"}, "index": 0, "type": "image_url"}] + ) images_chunk = _chunk(images_delta) assert _stream_chunks_have_generated_content([images_chunk]) is True - annotations_delta = Delta( - annotations=[{"type": "url_citation", "url_citation": {"url": "https://example.com"}}] - ) + annotations_delta = Delta(annotations=[{"type": "url_citation", "url_citation": {"url": "https://example.com"}}]) annotations_chunk = _chunk(annotations_delta) assert _stream_chunks_have_generated_content([annotations_chunk]) is True @@ -10120,12 +9793,8 @@ def test_get_configured_token_limits_skips_wildcard_pattern_matching(): ] ) - with patch.object( - router.pattern_router, "route", side_effect=AssertionError("pattern route called") - ): - assert router.get_configured_token_limits( - "bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0" - ) == (None, None) + with patch.object(router.pattern_router, "route", side_effect=AssertionError("pattern route called")): + assert router.get_configured_token_limits("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") == (None, None) def test_get_configured_token_limits_treats_malformed_values_as_absent(): @@ -10348,13 +10017,8 @@ def test_get_model_listing_info_skips_wildcard_pattern_matching(): ] ) - with patch.object( - router.pattern_router, "route", side_effect=AssertionError("pattern route called") - ): - assert ( - router.get_model_listing_info("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") - is None - ) + with patch.object(router.pattern_router, "route", side_effect=AssertionError("pattern route called")): + assert router.get_model_listing_info("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") is None def test_get_configured_mode_reads_deployment_model_info(): @@ -10396,13 +10060,8 @@ def test_get_configured_mode_skips_wildcard_pattern_matching(): ] ) - with patch.object( - router.pattern_router, "route", side_effect=AssertionError("pattern route called") - ): - assert ( - router.get_configured_mode("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") - is None - ) + with patch.object(router.pattern_router, "route", side_effect=AssertionError("pattern route called")): + assert router.get_configured_mode("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") is None def test_get_configured_mode_treats_malformed_values_as_absent(): @@ -10461,13 +10120,8 @@ def test_get_configured_display_name_skips_wildcard_pattern_matching(): ] ) - with patch.object( - router.pattern_router, "route", side_effect=AssertionError("pattern route called") - ): - assert ( - router.get_configured_display_name("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") - is None - ) + with patch.object(router.pattern_router, "route", side_effect=AssertionError("pattern route called")): + assert router.get_configured_display_name("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0") is None @pytest.mark.parametrize( @@ -10496,7 +10150,10 @@ def test_get_configured_service_tiers_returns_one_value_per_deployment_in_model_ "litellm_params": {"model": "openai/gpt-6-astra"}, "model_info": {"service_tiers": ["ultrafast"]}, }, - {"model_name": "gpt-6-astra", "litellm_params": {"model": "openai/gpt-6-astra", "api_base": "https://a.example"}}, + { + "model_name": "gpt-6-astra", + "litellm_params": {"model": "openai/gpt-6-astra", "api_base": "https://a.example"}, + }, { "model_name": "gpt-6-astra", "litellm_params": {"model": "openai/gpt-6-astra", "api_base": "https://b.example"}, @@ -10894,13 +10551,16 @@ async def test_acreate_batch_request_bedrock_tags_override_deployment_tags(): mock_client = MagicMock() mock_client.post = AsyncMock(side_effect=lambda *args, **kwargs: fake_response()) - with patch.object( - CommonBatchFilesUtils, - "sign_aws_request", - return_value=({"Authorization": "signed"}, b"{}"), - ) as mock_sign, patch( - "litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client", - return_value=mock_client, + with ( + patch.object( + CommonBatchFilesUtils, + "sign_aws_request", + return_value=({"Authorization": "signed"}, b"{}"), + ) as mock_sign, + patch( + "litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client", + return_value=mock_client, + ), ): await router.acreate_batch( model="bedrock-batch-model", @@ -10929,13 +10589,13 @@ async def test_avector_store_search_injects_router(): """ from litellm.types.vector_stores import VectorStoreSearchResponse - expected_response = VectorStoreSearchResponse( - object="vector_store.search_results.page", search_query="q", data=[] - ) + expected_response = VectorStoreSearchResponse(object="vector_store.search_results.page", search_query="q", data=[]) mock_asearch = AsyncMock(return_value=expected_response) # Router.__init__ binds asearch via a local import, so patch the module # attribute before constructing the Router. - with patch("litellm.vector_stores.main.asearch", new=mock_asearch): # test-quality-ok: the SDK call is the only place the injected router kwarg is observable + with patch( + "litellm.vector_stores.main.asearch", new=mock_asearch + ): # test-quality-ok: the SDK call is the only place the injected router kwarg is observable router = litellm.Router( model_list=[ { @@ -10969,7 +10629,9 @@ async def test_avector_store_create_does_not_inject_router(): } ] ) - with patch("litellm.vector_stores.main.acreate", new=mock_acreate): # test-quality-ok: the SDK call is the only place a leaked router kwarg would surface + with patch( + "litellm.vector_stores.main.acreate", new=mock_acreate + ): # test-quality-ok: the SDK call is the only place a leaked router kwarg would surface create_response = await router.avector_store_create(model=None, custom_llm_provider="openai") assert create_response is expected_response @@ -10985,13 +10647,13 @@ def test_vector_store_search_injects_router(): """ from litellm.types.vector_stores import VectorStoreSearchResponse - expected_response = VectorStoreSearchResponse( - object="vector_store.search_results.page", search_query="q", data=[] - ) + expected_response = VectorStoreSearchResponse(object="vector_store.search_results.page", search_query="q", data=[]) mock_search = MagicMock(return_value=expected_response) # Router.__init__ binds search via a local import, so patch the module # attribute before constructing the Router. - with patch("litellm.vector_stores.main.search", new=mock_search): # test-quality-ok: the SDK call is the only place the injected router kwarg is observable + with patch( + "litellm.vector_stores.main.search", new=mock_search + ): # test-quality-ok: the SDK call is the only place the injected router kwarg is observable router = litellm.Router( model_list=[ { @@ -11000,9 +10662,7 @@ def test_vector_store_search_injects_router(): } ] ) - search_response = router.vector_store_search( - vector_store_id="v", query="q", custom_llm_provider="s3_vectors" - ) + search_response = router.vector_store_search(vector_store_id="v", query="q", custom_llm_provider="s3_vectors") assert search_response is expected_response mock_search.assert_called_once() @@ -11016,7 +10676,9 @@ def test_vector_store_create_does_not_inject_router(): mock_create = MagicMock(return_value=expected_response) # Router.__init__ binds create via a local import, so patch the module # attribute before constructing the Router. - with patch("litellm.vector_stores.main.create", new=mock_create): # test-quality-ok: the SDK call is the only place a leaked router kwarg would surface + with patch( + "litellm.vector_stores.main.create", new=mock_create + ): # test-quality-ok: the SDK call is the only place a leaked router kwarg would surface router = litellm.Router( model_list=[ { @@ -11051,9 +10713,7 @@ class TestPreRoutingStrategyRegistryLifecycle: def _complexity_router_params(default_model: str, tags=None) -> dict: return { "model": "auto_router/complexity_router", - "complexity_router_config": { - "tiers": {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o", "COMPLEX": "gpt-4o"} - }, + "complexity_router_config": {"tiers": {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o", "COMPLEX": "gpt-4o"}}, "complexity_router_default_model": default_model, **({"tags": tags} if tags else {}), } @@ -11358,9 +11018,7 @@ class TestPreRoutingStrategyRegistryLifecycle: deployment=Deployment( model_name="hybrid-router", litellm_params=LiteLLM_Params( - **self._hybrid_router_params( - {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o", "COMPLEX": "gpt-4o"} - ) + **self._hybrid_router_params({"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o", "COMPLEX": "gpt-4o"}) ), model_info=ModelInfo(id="router-1", db_model=True), ) @@ -11469,9 +11127,7 @@ class TestPreRoutingStrategyRegistryLifecycle: ({"model": "openai/gpt-4o"}, False), ] for params, expected in cases: - actual = router._deployment_participates_in_adaptive_routing( - litellm_params=LiteLLM_Params(**params) - ) + actual = router._deployment_participates_in_adaptive_routing(litellm_params=LiteLLM_Params(**params)) assert actual is expected, params["model"] @@ -11658,22 +11314,16 @@ class TestUpsertDeploymentRollback: router.delete_deployment(id="prod-1") assert router.has_model_id("prod-1") is False - router._restore_deployment_after_failed_upsert( - previous_deployment=previous, model_id="prod-1" - ) + router._restore_deployment_after_failed_upsert(previous_deployment=previous, model_id="prod-1") restored = router.get_deployment(model_id="prod-1") assert restored is not None assert restored.litellm_params.model == "gpt-4o" - router._restore_deployment_after_failed_upsert( - previous_deployment=previous, model_id="prod-1" - ) + router._restore_deployment_after_failed_upsert(previous_deployment=previous, model_id="prod-1") assert len(router.model_list) == 1 - router._restore_deployment_after_failed_upsert( - previous_deployment=None, model_id="prod-1" - ) + router._restore_deployment_after_failed_upsert(previous_deployment=None, model_id="prod-1") assert len(router.model_list) == 1 @@ -12433,18 +12083,14 @@ class TestAutoRouterSharedModelNameConnectionParams: return httpx.Response( status_code=200, json={ - "candidates": [ - {"content": {"parts": [{"text": "Paris"}], "role": "model"}, "finishReason": "STOP"} - ], + "candidates": [{"content": {"parts": [{"text": "Paris"}], "role": "model"}, "finishReason": "STOP"}], "usageMetadata": {"promptTokenCount": 5, "candidatesTokenCount": 1, "totalTokenCount": 6}, "modelVersion": "gemini-3.6-flash", }, request=httpx.Request("POST", "https://generativelanguage.googleapis.com"), ) - @pytest.mark.parametrize( - "plain_entry_first", [True, False], ids=["plain_entry_first", "marker_entry_first"] - ) + @pytest.mark.parametrize("plain_entry_first", [True, False], ids=["plain_entry_first", "marker_entry_first"]) async def test_routed_tier_call_goes_out_on_its_own_endpoint_and_credentials(self, plain_entry_first): """The outbound provider request for the routed tier hits the tier's own Gemini host with the tier's own key, never the plain sibling's api_base or api_key.""" @@ -12568,9 +12214,7 @@ async def _drive_cyclic_fallback(router, capture, recorder=None, **request_kwarg litellm.callbacks.append(recorder) try: with pytest.raises(litellm.InternalServerError): - await router.acompletion( - model="group-a", messages=[{"role": "user", "content": "hi"}], **request_kwargs - ) + await router.acompletion(model="group-a", messages=[{"role": "user", "content": "hi"}], **request_kwargs) finally: router_logger.removeHandler(capture) router_logger.setLevel(previous_level) @@ -12844,9 +12488,7 @@ async def test_fallback_failure_detail_from_upstream_is_bounded(): await _drive_cyclic_fallback( _cyclic_fallback_router(), capture, - mock_response=litellm.InternalServerError( - message=huge_message, llm_provider="openai", model="group-a" - ), + mock_response=litellm.InternalServerError(message=huge_message, llm_provider="openai", model="group-a"), ) assert capture.messages, "the fallback failure path did not log at ERROR" @@ -12912,9 +12554,7 @@ def test_ensure_deployment_affinity_callback_is_idempotent(): try: router._ensure_deployment_affinity_callback() router._ensure_deployment_affinity_callback() - affinity_callbacks = [ - cb for cb in router.optional_callbacks or [] if isinstance(cb, DeploymentAffinityCheck) - ] + affinity_callbacks = [cb for cb in router.optional_callbacks or [] if isinstance(cb, DeploymentAffinityCheck)] assert len(affinity_callbacks) == 1 finally: for cb in router.optional_callbacks or []: @@ -13056,9 +12696,7 @@ class TestModelGroupAliasReachesPreRoutingStrategies: router = self._router("auto_routers") metadata: dict = {} - response = await router.acompletion( - model="smart-alias", messages=self._messages(), metadata=metadata - ) + response = await router.acompletion(model="smart-alias", messages=self._messages(), metadata=metadata) assert response.choices[0].message.content == "routed by the tier" assert metadata["model_group"] == "smart-alias" @@ -13300,6 +12938,7 @@ class TestTeamPublicNameReachesPreRoutingStrategies: assert response is not None assert response.model == "gemini-flash" + @pytest.mark.asyncio async def test_strategy_resolution_agrees_with_the_deployment_path_for_every_principal(self): router = self._router( @@ -13358,7 +12997,6 @@ class TestTeamPublicNameReachesPreRoutingStrategies: with pytest.raises(litellm.BadRequestError, match="multiple teams"): two_teams._team_deployments_across_teams(self.PUBLIC_NAME) - def test_compression_policy_follows_the_same_resolution_for_every_principal(self): from litellm.proxy.guardrails.auto_router_compression import AutoRouterCompressionPolicy, policy_for_model @@ -13582,7 +13220,6 @@ class TestAutoRouterCompressionDecoupling: @pytest.mark.usefixtures("local_model_cost_map") - @pytest.mark.usefixtures("local_model_cost_map") class TestAzureBaseModelFallbackLogging: """When an azure deployment has no base_model but its model name is a known @@ -13609,17 +13246,14 @@ class TestAzureBaseModelFallbackLogging: def test_map_known_deployment_name_resolves_without_error_log(self): router = self._router_with_azure_deployment("azure/gpt-4o") - with patch( - "litellm.router.verbose_router_logger.error" - ) as mock_error: + with patch("litellm.router.verbose_router_logger.error") as mock_error: model_info = router.get_router_model_info( deployment=None, received_model_name="my-group", id="azure-base-model-test-id" ) - assert not any( - "Could not identify azure model" in str(call) - for call in mock_error.call_args_list - ), f"unexpected error log: {mock_error.call_args_list}" + assert not any("Could not identify azure model" in str(call) for call in mock_error.call_args_list), ( + f"unexpected error log: {mock_error.call_args_list}" + ) # the fallback resolution must actually surface the map values assert model_info["max_input_tokens"] == litellm.model_cost["azure/gpt-4o"]["max_input_tokens"] assert model_info["input_cost_per_token"] == litellm.model_cost["azure/gpt-4o"]["input_cost_per_token"] @@ -13627,17 +13261,14 @@ class TestAzureBaseModelFallbackLogging: def test_unmappable_deployment_name_still_logs_error(self): router = self._router_with_azure_deployment("azure/my-custom-deployment-name") - with patch( - "litellm.router.verbose_router_logger.error" - ) as mock_error: + with patch("litellm.router.verbose_router_logger.error") as mock_error: model_info = router.get_router_model_info( deployment=None, received_model_name="my-group", id="azure-base-model-test-id" ) - assert any( - "Could not identify azure model" in str(call) - for call in mock_error.call_args_list - ), "expected the error log for an unmappable azure deployment name" + assert any("Could not identify azure model" in str(call) for call in mock_error.call_args_list), ( + "expected the error log for an unmappable azure deployment name" + ) # unmappable names resolve to a zeroed stub — unchanged behavior assert model_info.get("max_input_tokens") is None @@ -13664,6 +13295,7 @@ class TestAzureBaseModelFallbackLogging: ) assert model_info["max_input_tokens"] == litellm.model_cost["azure/gpt-4o-mini"]["max_input_tokens"] + def test_model_group_info_intersects_supported_reasoning_efforts(): router = litellm.Router( model_list=[ @@ -13754,7 +13386,6 @@ def test_model_group_info_reasoning_efforts_are_unknown_when_any_deployment_is_o assert result.supported_reasoning_efforts is None - @pytest.mark.parametrize( "model,provider,expected", [ @@ -13774,11 +13405,15 @@ def test_model_group_info_reasoning_efforts_are_unknown_when_any_deployment_is_o def test_model_group_info_fast_mode_uses_exact_provider_catalog( local_model_cost_map: None, model: str, provider: str | None, expected: bool, operator_flag: bool ) -> None: - router: Final = Router(model_list=[{ - "model_name": "fast-group", - "litellm_params": {"model": model, "custom_llm_provider": provider, "api_key": "fake-key"}, - "model_info": {"supports_fast_mode": operator_flag}, - }]) + router: Final = Router( + model_list=[ + { + "model_name": "fast-group", + "litellm_params": {"model": model, "custom_llm_provider": provider, "api_key": "fake-key"}, + "model_info": {"supports_fast_mode": operator_flag}, + } + ] + ) result: Final = router.get_model_group_info("fast-group") @@ -13790,16 +13425,21 @@ def test_model_group_info_fast_mode_uses_exact_provider_catalog( def test_model_group_info_fast_mode_fails_closed_without_explicit_boolean( local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch, flag: object ) -> None: - entry: Final = {key: value for key, value in litellm.model_cost["claude-opus-5"].items() - if key != "supports_fast_mode"} + entry: Final = { + key: value for key, value in litellm.model_cost["claude-opus-5"].items() if key != "supports_fast_mode" + } if flag is not None: entry["supports_fast_mode"] = flag monkeypatch.setitem(litellm.model_cost, "claude-opus-5", entry) - router: Final = Router(model_list=[{ - "model_name": "fast-group", - "litellm_params": {"model": "anthropic/claude-opus-5", "api_key": "fake-key"}, - "model_info": {"supports_fast_mode": True}, - }]) + router: Final = Router( + model_list=[ + { + "model_name": "fast-group", + "litellm_params": {"model": "anthropic/claude-opus-5", "api_key": "fake-key"}, + "model_info": {"supports_fast_mode": True}, + } + ] + ) result: Final = router.get_model_group_info("fast-group") @@ -13807,24 +13447,30 @@ def test_model_group_info_fast_mode_fails_closed_without_explicit_boolean( assert result.supports_fast_mode is False -@pytest.mark.parametrize("other_model,expected", [ - ("anthropic/claude-opus-4-8", True), - ("anthropic/claude-opus-4-7", False), - ("anthropic/off-map-opus", False), - ("vertex_ai/claude-opus-5", False), - ("bedrock/claude-opus-5", False), -]) +@pytest.mark.parametrize( + "other_model,expected", + [ + ("anthropic/claude-opus-4-8", True), + ("anthropic/claude-opus-4-7", False), + ("anthropic/off-map-opus", False), + ("vertex_ai/claude-opus-5", False), + ("bedrock/claude-opus-5", False), + ], +) @pytest.mark.parametrize("reverse", [True, False]) def test_model_group_info_fast_mode_requires_every_deployment( local_model_cost_map: None, other_model: str, expected: bool, reverse: bool ) -> None: - models: Final = (other_model, "anthropic/claude-opus-5") if reverse else ( - "anthropic/claude-opus-5", other_model + models: Final = (other_model, "anthropic/claude-opus-5") if reverse else ("anthropic/claude-opus-5", other_model) + router: Final = Router( + model_list=[ + { + "model_name": "fast-group", + "litellm_params": {"model": model, "api_key": "fake-key"}, + } + for model in models + ] ) - router: Final = Router(model_list=[{ - "model_name": "fast-group", - "litellm_params": {"model": model, "api_key": "fake-key"}, - } for model in models]) result: Final = router.get_model_group_info("fast-group") @@ -14064,6 +13710,7 @@ class TestAddDeploymentApiBaseProviderResolution: assert deployment is not None assert deployment.litellm_params.custom_llm_provider == "openai" + # ===================================================================== # anthropic_messages mid-stream-fallback helpers, added for #24004 # (mid-stream fallback not supported for anthropic_messages route type). @@ -14178,10 +13825,7 @@ class _AnthropicMessagesFallbackByteStream: def _anthropic_messages_overloaded_error_chunk() -> bytes: - return ( - b"event: error\n" - b'data: {"type": "error", "error": {"type": "overloaded_error", "message": "Overloaded"}}\n\n' - ) + return b'event: error\ndata: {"type": "error", "error": {"type": "overloaded_error", "message": "Overloaded"}}\n\n' def _anthropic_messages_invalid_request_error_chunk() -> bytes: @@ -14226,9 +13870,7 @@ async def test_anthropic_messages_streaming_iterator_passthrough(): [_anthropic_messages_content_chunk("hi"), _anthropic_messages_content_chunk(" there")] ) - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, initial_kwargs={"model": "primary"} - ) + wrapped = await router._aanthropic_messages_streaming_iterator(response=source, initial_kwargs={"model": "primary"}) collected = [chunk async for chunk in wrapped] assert collected == [_anthropic_messages_content_chunk("hi"), _anthropic_messages_content_chunk(" there")] @@ -14247,12 +13889,14 @@ async def test_anthropic_messages_streaming_iterator_flushes_buffered_lifecycle_ [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("hi"), message_stop] ) - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, initial_kwargs={"model": "primary"} - ) + wrapped = await router._aanthropic_messages_streaming_iterator(response=source, initial_kwargs={"model": "primary"}) collected = [chunk async for chunk in wrapped] - assert collected == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("hi"), message_stop] + assert collected == [ + _anthropic_messages_message_start_chunk(), + _anthropic_messages_content_chunk("hi"), + message_stop, + ] @pytest.mark.asyncio @@ -14264,9 +13908,7 @@ async def test_anthropic_messages_streaming_iterator_flushes_buffered_frames_on_ message_stop = b'event: message_stop\ndata: {"type": "message_stop"}\n\n' source = _AnthropicMessagesFakeByteStream([_anthropic_messages_message_start_chunk(), message_stop]) - wrapped = await router._aanthropic_messages_streaming_iterator( - response=source, initial_kwargs={"model": "primary"} - ) + wrapped = await router._aanthropic_messages_streaming_iterator(response=source, initial_kwargs={"model": "primary"}) collected = [chunk async for chunk in wrapped] assert collected == [_anthropic_messages_message_start_chunk(), message_stop] @@ -14317,7 +13959,9 @@ async def test_anthropic_messages_ping_behind_buffered_lifecycle_frame_is_forwar await content_released.wait() yield _anthropic_messages_content_chunk("hi") - wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source(), initial_kwargs={"model": "primary"} + ) assert await asyncio.wait_for(wrapped.__anext__(), timeout=1) == _anthropic_messages_ping_chunk() content_released.set() @@ -14362,7 +14006,9 @@ async def test_anthropic_messages_no_fallback_message_start_reaches_client_befor await content_released.wait() yield _anthropic_messages_content_chunk("hi") - wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source(), initial_kwargs={"model": "primary"} + ) assert await asyncio.wait_for(wrapped.__anext__(), timeout=1) == _anthropic_messages_message_start_chunk() content_released.set() @@ -14427,7 +14073,9 @@ async def test_anthropic_messages_default_wildcard_fallback_still_buffers_lifecy await content_released.wait() yield _anthropic_messages_content_chunk("hi") - wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source(), initial_kwargs={"model": "primary"} + ) pending = asyncio.ensure_future(wrapped.__anext__()) await asyncio.sleep(0.2) @@ -14467,8 +14115,15 @@ def _anthropic_messages_two_order_primary_model_list() -> list: id="wildcard-overridden-by-request-none", ), pytest.param({"fallbacks": [{"*": ["fallback"]}]}, {"model": "primary"}, True, id="wildcard"), - pytest.param({"fallbacks": None}, {"model": "primary", "fallbacks": [{"model": "fallback"}]}, True, id="request-dict-fallback"), - pytest.param({"fallbacks": None}, {"model": "primary", "fallbacks": ["fallback"]}, True, id="request-list-fallback"), + pytest.param( + {"fallbacks": None}, + {"model": "primary", "fallbacks": [{"model": "fallback"}]}, + True, + id="request-dict-fallback", + ), + pytest.param( + {"fallbacks": None}, {"model": "primary", "fallbacks": ["fallback"]}, True, id="request-list-fallback" + ), pytest.param( {"fallbacks": [{"primary": ["fallback"]}]}, {"model": "primary", "disable_fallbacks": True}, @@ -14481,7 +14136,9 @@ def _anthropic_messages_two_order_primary_model_list() -> list: True, id="content-policy-fallback", ), - pytest.param({"fallbacks": None, "enable_weighted_failover": True}, {"model": "primary"}, True, id="weighted-failover"), + pytest.param( + {"fallbacks": None, "enable_weighted_failover": True}, {"model": "primary"}, True, id="weighted-failover" + ), ], ) def test_anthropic_messages_stream_can_fall_back_direct_call(router_kwargs, request_kwargs, expected): @@ -14554,7 +14211,9 @@ async def test_anthropic_messages_order_fallback_still_buffers_lifecycle_frames( await content_released.wait() yield _anthropic_messages_content_chunk("hi") - wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source(), initial_kwargs={"model": "primary"} + ) pending = asyncio.ensure_future(wrapped.__anext__()) await asyncio.sleep(0.2) @@ -14599,7 +14258,9 @@ async def test_anthropic_messages_leading_ping_keepalive_is_forwarded_live(): yield _anthropic_messages_message_start_chunk() yield _anthropic_messages_content_chunk("hi") - wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source(), initial_kwargs={"model": "primary"} + ) assert await asyncio.wait_for(wrapped.__anext__(), timeout=1) == _anthropic_messages_ping_chunk() content_released.set() @@ -15296,7 +14957,9 @@ class _AnthropicMessagesScriptedProvider: async def __call__(self, **kwargs): litellm_metadata = kwargs.get("litellm_metadata") or {} - self.calls.append((kwargs["model"], litellm_metadata.get("attempted_retries"), litellm_metadata.get("max_retries"))) + self.calls.append( + (kwargs["model"], litellm_metadata.get("attempted_retries"), litellm_metadata.get("max_retries")) + ) assert self._streams, "provider called more times than scripted" return self._streams.pop(0)() @@ -15314,7 +14977,9 @@ def _anthropic_messages_transport_drop(original_exception: Exception | None = No def _anthropic_messages_dropped_before_content(): - return _AnthropicMessagesRaisingByteStream([_anthropic_messages_message_start_chunk()], _anthropic_messages_transport_drop()) + return _AnthropicMessagesRaisingByteStream( + [_anthropic_messages_message_start_chunk()], _anthropic_messages_transport_drop() + ) def _anthropic_messages_bridge_error_chunk() -> bytes: @@ -15919,20 +15584,28 @@ async def test_anthropic_messages_retries_keep_the_budget_the_first_drop_committ }, { "model_name": "glm", - "litellm_params": {"model": "anthropic/glm-b", "api_key": "sk-test", "num_retries": sibling_num_retries}, + "litellm_params": { + "model": "anthropic/glm-b", + "api_key": "sk-test", + "num_retries": sibling_num_retries, + }, }, ], num_retries=0, fallbacks=None, ) router.set_custom_routing_strategy(_AnthropicMessagesAlternatingDeployments(router, "glm")) - provider = _AnthropicMessagesScriptedProvider(*[_anthropic_messages_dropped_before_content] * len(expected_counters)) + provider = _AnthropicMessagesScriptedProvider( + *[_anthropic_messages_dropped_before_content] * len(expected_counters) + ) stream = await _anthropic_messages_stream_through_router(router, provider) with pytest.raises(litellm.APIConnectionError): [chunk async for chunk in stream] - assert [model for model, _, _ in provider.calls] == (["anthropic/glm-a", "anthropic/glm-b"] * 2)[: len(expected_counters)] + assert [model for model, _, _ in provider.calls] == (["anthropic/glm-a", "anthropic/glm-b"] * 2)[ + : len(expected_counters) + ] assert [(attempted, budget) for _, attempted, budget in provider.calls] == expected_counters @@ -17419,9 +17092,9 @@ class TestTierParamsTheTargetAccepts: def test_declared_param_allowlist_ignores_malformed_declarations(self): """A str is iterable, so without the type guard a YAML scalar mistake like allowed_openai_params: reasoning_effort would allowlist single characters.""" - assert litellm.Router._declared_param_allowlist({"allowed_openai_params": ["reasoning_effort", 3]}) == frozenset( - {"reasoning_effort"} - ) + assert litellm.Router._declared_param_allowlist( + {"allowed_openai_params": ["reasoning_effort", 3]} + ) == frozenset({"reasoning_effort"}) assert litellm.Router._declared_param_allowlist({"allowed_openai_params": "reasoning_effort"}) == frozenset() assert litellm.Router._declared_param_allowlist({}) == frozenset() @@ -17501,7 +17174,11 @@ class TestTierParamsTheTargetAccepts: @pytest.mark.parametrize( "deployment", - [{"model_name": "x"}, {"model_name": "x", "litellm_params": {}}, {"model_name": "x", "litellm_params": {"model": "not-a-real-provider/nope"}}], + [ + {"model_name": "x"}, + {"model_name": "x", "litellm_params": {}}, + {"model_name": "x", "litellm_params": {"model": "not-a-real-provider/nope"}}, + ], ) def test_deployment_accepts_param_fails_open(self, deployment): """An unresolvable deployment must not be the reason a param is dropped.""" @@ -17806,9 +17483,7 @@ class TestPreRoutingTierDrivesFallbacks: async def test_the_selected_tier_fallback_chain_runs(self): router = self._router([{"tier1": ["backup-a"]}]) - response = await router.acompletion( - model="smart-router", messages=[{"role": "user", "content": "hi"}] - ) + response = await router.acompletion(model="smart-router", messages=[{"role": "user", "content": "hi"}]) assert response.choices[0].message.content == "from backup-a" @@ -17841,9 +17516,7 @@ class TestPreRoutingTierDrivesFallbacks: async def test_a_request_without_a_pre_routing_hook_still_uses_its_own_group(self): router = self._router([{"tier1": ["backup-a"]}]) - response = await router.acompletion( - model="tier1", messages=[{"role": "user", "content": "hi"}] - ) + response = await router.acompletion(model="tier1", messages=[{"role": "user", "content": "hi"}]) assert response.choices[0].message.content == "from backup-a" @@ -19089,7 +18762,9 @@ async def test_router_deployment_slot_rejects_while_held_and_frees_slot_on_exit( async with router._deployment_slot(deployment=deployment.model_dump(), kwargs=kwargs, parent_otel_span=None): with pytest.raises(litellm.RateLimitError) as overflow: - async with router._deployment_slot(deployment=deployment.model_dump(), kwargs=kwargs, parent_otel_span=None): + async with router._deployment_slot( + deployment=deployment.model_dump(), kwargs=kwargs, parent_otel_span=None + ): pass assert overflow.value.status_code == 429 assert "slot-deployment" in overflow.value.message @@ -19324,18 +18999,25 @@ class TestMemberAutoRouterInference: self.cache = UserApiKeyCache() self.team = LiteLLM_TeamTable( - team_id="router-team", models=["member-router", "permitted-model"], + team_id="router-team", + models=["member-router", "permitted-model"], members_with_roles=[Member(user_id="router-member", role="user")], ) self.actor = UserAPIKeyAuth( - user_id="router-member", team_id="router-team", user_role=LitellmUserRoles.INTERNAL_USER, - models=["member-router", "permitted-model"], api_key="test-key-hash", config={"timeout": 60}, + user_id="router-member", + team_id="router-team", + user_role=LitellmUserRoles.INTERNAL_USER, + models=["member-router", "permitted-model"], + api_key="test-key-hash", + config={"timeout": 60}, + ) + self.database = SimpleNamespace( + db=SimpleNamespace( + litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=self.team)), + litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(return_value=None)), + litellm_accessgrouptable=SimpleNamespace(find_unique=AsyncMock()), + ) ) - self.database = SimpleNamespace(db=SimpleNamespace( - litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=self.team)), - litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(return_value=None)), - litellm_accessgrouptable=SimpleNamespace(find_unique=AsyncMock()), - )) monkeypatch.setattr(proxy_server, "user_api_key_cache", self.cache) monkeypatch.setattr(proxy_server, "prisma_client", self.database) @@ -19345,44 +19027,71 @@ class TestMemberAutoRouterInference: return { "model_name": "model_name_router-team_member-router", "litellm_params": { - "model": "auto_router/complexity_router", "complexity_router_default_model": target, + "model": "auto_router/complexity_router", + "complexity_router_default_model": target, "complexity_router_config": { - "tiers": dict.fromkeys(("SIMPLE", "MEDIUM", "COMPLEX", "REASONING"), target), "adaptive": False, + "tiers": dict.fromkeys(("SIMPLE", "MEDIUM", "COMPLEX", "REASONING"), target), + "adaptive": False, **({"classifier_type": "llm", "classifier_llm_config": {"model": target}} if classifier else {}), }, - "tags": ["member" if member else "admin"], "timeout": 13.0 if member else 29.0, + "tags": ["member" if member else "admin"], + "timeout": 13.0 if member else 29.0, }, "model_info": { - "team_id": "router-team", "team_public_model_name": "member-router", "member_auto_router": member, + "team_id": "router-team", + "team_public_model_name": "member-router", + "member_auto_router": member, }, } @classmethod def _router(cls, *markers: dict[str, object]) -> Router: - return Router(model_list=[ - *(markers or (cls._marker(),)), - {"model_name": "permitted-model", "litellm_params": { - "model": "openai/gpt-4o-mini", "api_key": "test-key", "api_base": "https://api.openai.com/v1", - }}, - {"model_name": "restricted-model", "litellm_params": {"model": "openai/gpt-4o", "api_key": "test-key"}}, - ]) + return Router( + model_list=[ + *(markers or (cls._marker(),)), + { + "model_name": "permitted-model", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "test-key", + "api_base": "https://api.openai.com/v1", + }, + }, + {"model_name": "restricted-model", "litellm_params": {"model": "openai/gpt-4o", "api_key": "test-key"}}, + ] + ) def _request( - self, *, actor: UserAPIKeyAuth | None = None, metadata_name: str = "metadata", tag: str = "member", + self, + *, + actor: UserAPIKeyAuth | None = None, + metadata_name: str = "metadata", + tag: str = "member", ) -> dict[str, object]: from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup return LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata( - data={metadata_name: {"tags": [tag]}, **({"metadata": {"user_api_key_auth": {"user_role": "proxy_admin"}}} - if metadata_name == "litellm_metadata" else {})}, - user_api_key_dict=actor or self.actor, _metadata_variable_name=metadata_name, + data={ + metadata_name: {"tags": [tag]}, + **( + {"metadata": {"user_api_key_auth": {"user_role": "proxy_admin"}}} + if metadata_name == "litellm_metadata" + else {} + ), + }, + user_api_key_dict=actor or self.actor, + _metadata_variable_name=metadata_name, ) async def _route( - self, router: Router, request: dict[str, object] | None = None, model: str = "member-router", + self, + router: Router, + request: dict[str, object] | None = None, + model: str = "member-router", ) -> PreRoutingHookResponse: response: Final = await router.async_pre_routing_hook( - model=model, request_kwargs=request if request is not None else self._request(), + model=model, + request_kwargs=request if request is not None else self._request(), messages=[{"role": "user", "content": "Hello"}], ) assert response is not None @@ -19391,36 +19100,68 @@ class TestMemberAutoRouterInference: @pytest.mark.asyncio @pytest.mark.parametrize("metadata_name", ("metadata", "litellm_metadata")) async def test_cached_roster_revocation_blocks_classifier_and_session_rebinding( - self, metadata_name: str, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, + self, + metadata_name: str, + respx_mock: respx.MockRouter, + monkeypatch: pytest.MonkeyPatch, ) -> None: from litellm.proxy.auth.auth_checks import delete_cache_team_object monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") router: Final = self._router(self._marker(classifier=True)) - classify: Final = respx_mock.post("https://api.openai.com/v1/chat/completions").respond(200, json={ - "id": "classifier", "object": "chat.completion", "created": 0, "model": "gpt-4o-mini", - "choices": [{"index": 0, "message": {"content": '{"tier":"SIMPLE"}', "role": "assistant"}, - "finish_reason": "stop"}], - "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, - }) - request: Final = {**self._request(metadata_name=metadata_name), "proxy_server_request": {"headers": { - "x-claude-code-session-id": "member-router-session", "x-app": "cli", - }}} + classify: Final = respx_mock.post("https://api.openai.com/v1/chat/completions").respond( + 200, + json={ + "id": "classifier", + "object": "chat.completion", + "created": 0, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"content": '{"tier":"SIMPLE"}', "role": "assistant"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + ) + request: Final = { + **self._request(metadata_name=metadata_name), + "proxy_server_request": { + "headers": { + "x-claude-code-session-id": "member-router-session", + "x-app": "cli", + } + }, + } first: Final = await self._route(router, request) assert first.model == "permitted-model" and first.routing_decision is not None assert first.routing_decision["cause"] == "llm_classifier" assert (await self._route(router, request)).model == "permitted-model" assert self.database.db.litellm_teamtable.find_unique.await_count == 1 assert self.database.db.litellm_teammembership.find_unique.await_count == 1 - self.database.db.litellm_teamtable.find_unique.return_value = self.team.model_copy(update={"members_with_roles": []}) + self.database.db.litellm_teamtable.find_unique.return_value = self.team.model_copy( + update={"members_with_roles": []} + ) await delete_cache_team_object( - team_id=self.team.team_id, team_alias=None, user_api_key_cache=self.cache, proxy_logging_obj=None, + team_id=self.team.team_id, + team_alias=None, + user_api_key_cache=self.cache, + proxy_logging_obj=None, ) with pytest.raises(HTTPException, match="no longer a member"): await self._route(router, request) - rebound: Final = {**request, "proxy_server_request": {"headers": { - "x-claude-code-session-id": "member-router-session", "x-app": "cli", "x-claude-code-agent-id": "subagent", - }}} + rebound: Final = { + **request, + "proxy_server_request": { + "headers": { + "x-claude-code-session-id": "member-router-session", + "x-app": "cli", + "x-claude-code-agent-id": "subagent", + } + }, + } with pytest.raises(HTTPException, match="no longer a member"): await self._route(router, rebound, model="restricted-model") assert classify.call_count == 2 @@ -19430,11 +19171,23 @@ class TestMemberAutoRouterInference: async def test_member_router_fails_closed(self, state: str, monkeypatch: pytest.MonkeyPatch) -> None: from litellm.proxy import proxy_server - request: Final = {"metadata": {"user_api_key_team_id": "router-team", "user_api_key_auth": { - "team_id": "router-team", "user_role": "proxy_admin", - }}} if state == "forged" else self._request(actor=self.actor.model_copy( - update={"user_id": ""} if state == "empty-user" else {}, - )) + request: Final = ( + { + "metadata": { + "user_api_key_team_id": "router-team", + "user_api_key_auth": { + "team_id": "router-team", + "user_role": "proxy_admin", + }, + } + } + if state == "forged" + else self._request( + actor=self.actor.model_copy( + update={"user_id": ""} if state == "empty-user" else {}, + ) + ) + ) self.database.db.litellm_teamtable.find_unique.return_value = ( None if state == "deleted" else self.team.model_copy(update={"blocked": state == "blocked"}) ) @@ -19445,11 +19198,20 @@ class TestMemberAutoRouterInference: assert error.value.status_code == (503 if state == "unavailable" else 403) @pytest.mark.asyncio - @pytest.mark.parametrize("user_id,role", [(None, LitellmUserRoles.INTERNAL_USER), ("admin", LitellmUserRoles.PROXY_ADMIN)]) - async def test_service_key_and_admin_preserve_runtime_access(self, user_id: str | None, role: LitellmUserRoles) -> None: - assert (await self._route(self._router(), self._request( - actor=self.actor.model_copy(update={"user_id": user_id, "user_role": role}), - ))).model == "permitted-model" + @pytest.mark.parametrize( + "user_id,role", [(None, LitellmUserRoles.INTERNAL_USER), ("admin", LitellmUserRoles.PROXY_ADMIN)] + ) + async def test_service_key_and_admin_preserve_runtime_access( + self, user_id: str | None, role: LitellmUserRoles + ) -> None: + assert ( + await self._route( + self._router(), + self._request( + actor=self.actor.model_copy(update={"user_id": user_id, "user_role": role}), + ), + ) + ).model == "permitted-model" @pytest.mark.asyncio @pytest.mark.parametrize("ceiling", ("team", "key", "member", "organization", "project")) @@ -19460,35 +19222,56 @@ class TestMemberAutoRouterInference: from litellm.proxy._types import LiteLLM_ProjectTableCachedObj from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key - self.database.db.litellm_teamtable.find_unique.return_value = self.team.model_copy(update={ - "models": ["member-router"] if ceiling == "team" else self.team.models, - "organization_id": "router-org" if ceiling == "organization" else None, - }) + self.database.db.litellm_teamtable.find_unique.return_value = self.team.model_copy( + update={ + "models": ["member-router"] if ceiling == "team" else self.team.models, + "organization_id": "router-org" if ceiling == "organization" else None, + } + ) if ceiling == "member": await self.cache.async_set_cache( key=team_membership_reservation_cache_key(user_id="router-member", team_id="router-team"), - value=LiteLLM_TeamMembership(user_id="router-member", team_id="router-team", - litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["restricted-model"])), + value=LiteLLM_TeamMembership( + user_id="router-member", + team_id="router-team", + litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["restricted-model"]), + ), model_type=LiteLLM_TeamMembership, ) elif ceiling == "organization": await self.cache.async_set_cache( - key="org_id:router-org", value=LiteLLM_OrganizationTable( - organization_id="router-org", budget_id="org-budget", created_by="admin", updated_by="admin", + key="org_id:router-org", + value=LiteLLM_OrganizationTable( + organization_id="router-org", + budget_id="org-budget", + created_by="admin", + updated_by="admin", models=["restricted-model"], - ), model_type=LiteLLM_OrganizationTable, + ), + model_type=LiteLLM_OrganizationTable, ) elif ceiling == "project": await self.cache.async_set_cache( - key="project_id:router-project", value=LiteLLM_ProjectTableCachedObj( - project_id="router-project", team_id="router-team", models=["restricted-model"], - ), model_type=LiteLLM_ProjectTableCachedObj, + key="project_id:router-project", + value=LiteLLM_ProjectTableCachedObj( + project_id="router-project", + team_id="router-team", + models=["restricted-model"], + ), + model_type=LiteLLM_ProjectTableCachedObj, ) with pytest.raises(ProxyException, match="is not available for this API key"): - await self._route(self._router(), self._request(actor=self.actor.model_copy(update={ - "models": ["member-router"] if ceiling == "key" else self.actor.models, - "project_id": "router-project" if ceiling == "project" else None, - }))) + await self._route( + self._router(), + self._request( + actor=self.actor.model_copy( + update={ + "models": ["member-router"] if ceiling == "key" else self.actor.models, + "project_id": "router-project" if ceiling == "project" else None, + } + ) + ), + ) @pytest.mark.asyncio @pytest.mark.parametrize("group_owner", ("team", "key")) @@ -19496,22 +19279,32 @@ class TestMemberAutoRouterInference: from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast group: Final = LiteLLM_AccessGroupTable( - access_group_id="router-group", access_group_name="Router targets", access_model_names=["permitted-model"], + access_group_id="router-group", + access_group_name="Router targets", + access_model_names=["permitted-model"], ) self.database.db.litellm_accessgrouptable.find_unique.return_value = group - self.database.db.litellm_teamtable.find_unique.return_value = self.team.model_copy(update={ - "models": ["member-router"] if group_owner == "team" else self.team.models, - "access_group_ids": ["router-group"] if group_owner == "team" else [], - }) - request: Final = self._request(actor=self.actor.model_copy(update={ - "models": ["member-router"] if group_owner == "key" else self.actor.models, - "access_group_ids": ["router-group"] if group_owner == "key" else [], - })) + self.database.db.litellm_teamtable.find_unique.return_value = self.team.model_copy( + update={ + "models": ["member-router"] if group_owner == "team" else self.team.models, + "access_group_ids": ["router-group"] if group_owner == "team" else [], + } + ) + request: Final = self._request( + actor=self.actor.model_copy( + update={ + "models": ["member-router"] if group_owner == "key" else self.actor.models, + "access_group_ids": ["router-group"] if group_owner == "key" else [], + } + ) + ) router: Final = self._router() assert (await self._route(router, request)).model == "permitted-model" assert (await self._route(router, request)).model == "permitted-model" assert self.database.db.litellm_accessgrouptable.find_unique.await_count == 1 - self.database.db.litellm_accessgrouptable.find_unique.return_value = group.model_copy(update={"access_model_names": []}) + self.database.db.litellm_accessgrouptable.find_unique.return_value = group.model_copy( + update={"access_model_names": []} + ) await evict_and_broadcast(cache_keys=("access_group_id:router-group",), user_api_key_cache=self.cache) with pytest.raises(ProxyException, match="is not available for this API key"): await self._route(router, request) @@ -19522,13 +19315,16 @@ class TestMemberAutoRouterInference: router: Final = self._router(self._marker(member=False), self._marker()) request: Final = self._request() selected: Final = router._selected_strategy_marker_deployment( - model="model_name_router-team_member-router", strategy_tags=("member",), request_kwargs=request, + model="model_name_router-team_member-router", + strategy_tags=("member",), + request_kwargs=request, ) assert selected is not None and selected["model_info"]["member_auto_router"] is True assert (await self._route(router, request)).model == "permitted-model" assert request["timeout"] == 13.0 await self.cache.async_set_cache( - key="team_id:router-team", model_type=LiteLLM_TeamTable, + key="team_id:router-team", + model_type=LiteLLM_TeamTable, value=self.team.model_copy(update={"models": ["member-router"]}), ) with pytest.raises(ProxyException, match="is not available for this API key"): @@ -19544,7 +19340,9 @@ class TestMemberAutoRouterInference: router: Final = self._router(self._marker(member=False)) monkeypatch.setitem(sys.modules, "fastapi", None) monkeypatch.delitem(sys.modules, "litellm.proxy.auth.auto_router_checks", raising=False) - assert (await self._route(router, {"metadata": {"user_api_key_team_id": "router-team"}})).model == "restricted-model" + assert ( + await self._route(router, {"metadata": {"user_api_key_team_id": "router-team"}}) + ).model == "restricted-model" def _access_window_offsets(start_hours: float, end_hours: float, team_ids: list) -> dict: @@ -19589,10 +19387,12 @@ def test_access_windows_hide_reserved_deployment_from_other_teams(): def test_access_windows_raise_when_only_reserved_deployments_remain(): - router = Router(model_list=_reserved_model_list( - windows_for_reserved=[_access_window_offsets(-1, 1, ["team-a"])], - windows_for_open=[_access_window_offsets(-1, 1, ["team-a"])], - )[:1]) + router = Router( + model_list=_reserved_model_list( + windows_for_reserved=[_access_window_offsets(-1, 1, ["team-a"])], + windows_for_open=[_access_window_offsets(-1, 1, ["team-a"])], + )[:1] + ) for request_kwargs in ({"metadata": {"user_api_key_team_id": "team-b"}}, {}): with pytest.raises(litellm.BadRequestError, match="reserved for another team"): router._common_checks_available_deployment(model="gpt-4o-ptu", request_kwargs=request_kwargs) @@ -19655,9 +19455,11 @@ def test_access_windows_inactive_window_leaves_deployments_available(): def test_access_windows_apply_when_calling_by_model_id(): - router = Router(model_list=_reserved_model_list( - windows_for_reserved=[_access_window_offsets(-1, 1, ["team-a"])], - )) + router = Router( + model_list=_reserved_model_list( + windows_for_reserved=[_access_window_offsets(-1, 1, ["team-a"])], + ) + ) with pytest.raises(litellm.BadRequestError, match="reserved for another team"): router._common_checks_available_deployment( model="reserved-deployment", @@ -19824,7 +19626,9 @@ async def test_bare_model_group_served_by_wildcard_deployment_uses_provider_pref @pytest.mark.asyncio -async def test_bare_model_group_served_by_wildcard_deployment_uses_provider_prefixed_context_window_fallback_key() -> None: +async def test_bare_model_group_served_by_wildcard_deployment_uses_provider_prefixed_context_window_fallback_key() -> ( + None +): """The context-window chain is keyed the same way the ordinary chain is, so a key spelled like the wildcard deployment ("anthropic/claude-sonnet-4-6") must catch the bare group's context-window error too.""" router = litellm.Router( @@ -19884,9 +19688,7 @@ def test_bare_model_group_served_by_wildcard_deployment_has_provider_prefixed_co [ GuardrailRaisedException(guardrail_name="chunk-scanner", message="blocked"), HTTPException(status_code=403, detail={"error": "blocked", "guardrail_name": "chunk-scanner"}), - ModifyResponseException( - message="blocked", model="primary", request_data={}, guardrail_name="chunk-scanner" - ), + ModifyResponseException(message="blocked", model="primary", request_data={}, guardrail_name="chunk-scanner"), ], ) async def test_a_guardrail_verdict_is_neither_retried_nor_fallen_back(verdict: Exception) -> None: @@ -20087,6 +19889,7 @@ async def test_failure_rpm_increment_declares_the_router_usage_key_family(): assert seen == ["router_usage"] assert current_service_target() is None + class _SpanRecordingInMemoryCache(InMemoryCache): """Records the live OTel span each read runs under, so the test sees what a Redis span would nest in.""" diff --git a/tests/unit/test_router_block_helpers.py b/tests/unit/test_router_block_helpers.py index 443209bfe2f..c89d514b166 100644 --- a/tests/unit/test_router_block_helpers.py +++ b/tests/unit/test_router_block_helpers.py @@ -19,7 +19,7 @@ class TestAreAllDeploymentsBlocked: def test_all_blocked_returns_true(self): router = _make_router("gpt-4o", blocked=True) deployments = router.get_model_list(model_name="gpt-4o") or [] - assert router._are_all_deployments_blocked(deployments) is True + assert router.are_all_deployments_blocked(deployments) is True def test_one_not_blocked_returns_false(self): router = Router( @@ -40,11 +40,11 @@ class TestAreAllDeploymentsBlocked: ] ) deployments = router.get_model_list(model_name="gpt-4o") or [] - assert router._are_all_deployments_blocked(deployments) is False + assert router.are_all_deployments_blocked(deployments) is False def test_empty_list_returns_false(self): router = _make_router("gpt-4o") - assert router._are_all_deployments_blocked([]) is False + assert router.are_all_deployments_blocked([]) is False class TestIsModelFullyBlocked: diff --git a/tests/unit/test_router_model_cost_isolation.py b/tests/unit/test_router_model_cost_isolation.py index ff8c91cae70..4aed84d05c8 100644 --- a/tests/unit/test_router_model_cost_isolation.py +++ b/tests/unit/test_router_model_cost_isolation.py @@ -2220,7 +2220,7 @@ def test_replay_model_cost_registrations_survives_a_malformed_deployment(): litellm.model_cost = {"gpt-4o": {"litellm_provider": "openai", "mode": "chat"}} _invalidate_model_cost_lowercase_map() - router._replay_model_cost_registrations() + router.replay_model_cost_registrations() assert litellm.model_cost["healthy-id"]["max_input_tokens"] == 777 finally: diff --git a/tests/unit/test_router_order_fallback.py b/tests/unit/test_router_order_fallback.py index 93895bbde08..a86338c625a 100644 --- a/tests/unit/test_router_order_fallback.py +++ b/tests/unit/test_router_order_fallback.py @@ -18,7 +18,7 @@ from litellm import Router from litellm.integrations.custom_logger import CustomLogger from litellm.router_utils.prompt_caching_cache import PromptCachingCache from litellm.types.router import RouterRateLimitError -from litellm.utils import _get_deployment_order, get_order_filtered_deployments +from litellm.utils import get_deployment_order, get_order_filtered_deployments # --------------------------------------------------------------------------- # Unit tests for get_order_filtered_deployments @@ -94,6 +94,10 @@ class TestGetOrderFilteredDeployments: assert len(result) == 2 +def test_get_deployment_order_returns_unvalidated_order(): + assert get_deployment_order({"litellm_params": {"order": "first"}}) == "first" + + # --------------------------------------------------------------------------- # Integration tests for order-based fallback in Router # --------------------------------------------------------------------------- @@ -428,7 +432,7 @@ async def test_router_order_fallback_does_not_reselect_order_1_when_order_2_is_f async def async_filter_deployments( self, model, healthy_deployments, messages, request_kwargs=None, parent_otel_span=None ): - return [d for d in healthy_deployments if _get_deployment_order(d) != 2] + return [d for d in healthy_deployments if get_deployment_order(d) != 2] drop_order_2: Final = _DropOrder2() router = Router( @@ -643,18 +647,18 @@ async def test_text_completion_order_fallback_hop_does_not_send_target_order_ups def test_check_non_standard_fallback_format(): from litellm.router_utils.fallback_event_handlers import ( - _check_non_standard_fallback_format, + check_non_standard_fallback_format, ) # Standard formats - assert _check_non_standard_fallback_format([{"gpt-3.5-turbo": ["claude-3-haiku"]}]) == False - assert _check_non_standard_fallback_format([{"model": ["qwen-backup"]}]) == False - assert _check_non_standard_fallback_format([{"model": ["qwen-backup"], "region": ["us-east-1"]}]) == False + assert check_non_standard_fallback_format([{"gpt-3.5-turbo": ["claude-3-haiku"]}]) == False + assert check_non_standard_fallback_format([{"model": ["qwen-backup"]}]) == False + assert check_non_standard_fallback_format([{"model": ["qwen-backup"], "region": ["us-east-1"]}]) == False # Non-standard formats - assert _check_non_standard_fallback_format([{"model": "qwen-backup"}]) == True + assert check_non_standard_fallback_format([{"model": "qwen-backup"}]) == True assert ( - _check_non_standard_fallback_format([{"model": "qwen-backup", "messages": [{"role": "user", "content": "hi"}]}]) + check_non_standard_fallback_format([{"model": "qwen-backup", "messages": [{"role": "user", "content": "hi"}]}]) == True ) - assert _check_non_standard_fallback_format([{"model": ["qwen-backup"], "api_key": "some-key"}]) == True + assert check_non_standard_fallback_format([{"model": ["qwen-backup"], "api_key": "some-key"}]) == True diff --git a/tests/unit/test_router_weighted_failover.py b/tests/unit/test_router_weighted_failover.py index 9f05654f23f..9cb38b00e1e 100644 --- a/tests/unit/test_router_weighted_failover.py +++ b/tests/unit/test_router_weighted_failover.py @@ -810,7 +810,7 @@ async def test_maybe_run_weighted_failover_skips_when_remaining_all_in_cooldown( # Patch cooldown so B and C appear in cooldown. with patch( - "litellm.router._async_get_cooldown_deployments", + "litellm.router.async_get_cooldown_deployments", new=AsyncMock(return_value=["B", "C"]), ): result = await router._maybe_run_weighted_failover( @@ -869,7 +869,7 @@ async def test_maybe_run_weighted_failover_proceeds_when_one_healthy_remains( monkeypatch.setattr("litellm.router.run_async_fallback", _stub_run_async_fallback) with patch( - "litellm.router._async_get_cooldown_deployments", + "litellm.router.async_get_cooldown_deployments", new=AsyncMock(return_value=["B"]), ): result = await router._maybe_run_weighted_failover( @@ -935,7 +935,7 @@ async def test_failover_falls_through_to_external_fallback_when_remaining_in_coo # Put B in cooldown so weighted failover can't use it after A fails. with patch( - "litellm.router._async_get_cooldown_deployments", + "litellm.router.async_get_cooldown_deployments", new=AsyncMock(return_value=["B"]), ): response = await router.acompletion( diff --git a/tests/unit/test_secret_redaction.py b/tests/unit/test_secret_redaction.py index 71369c33a6b..1c1cfe56754 100644 --- a/tests/unit/test_secret_redaction.py +++ b/tests/unit/test_secret_redaction.py @@ -12,7 +12,7 @@ import pytest from litellm._logging import ( JsonFormatter, - _redact_string, + redact_string, _secret_filter, redact_internal_details_from_client_message, verbose_logger, @@ -21,7 +21,6 @@ from litellm._logging import ( ) from litellm.litellm_core_utils.secret_redaction import ( redact_internal_details, - redact_string, redact_structured_value, ) @@ -474,7 +473,7 @@ def test_vertex_error_message_no_credential_leak(): "Ensure the JSON is valid (check for unescaped newlines in private_key). " "Parse error: JSONDecodeError" ) - result = _redact_string(new_msg) + result = redact_string(new_msg) assert result == new_msg # nothing to redact diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 44b308a7e1a..783710e46f0 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -68,7 +68,7 @@ from litellm.utils import ( ProviderConfigManager, TextCompletionStreamWrapper, _check_provider_match, - _get_potential_model_names, + get_potential_model_names, _is_litellm_router_call, _is_streaming_request, _run_success_deployment_hook_on_converted_chat_stream, @@ -84,6 +84,13 @@ from litellm.utils import ( is_prompt_caching_valid_prompt, ) + +def test_get_base_model_from_metadata_returns_unvalidated_root_value(): + from litellm.utils import get_base_model_from_metadata + + assert get_base_model_from_metadata({"litellm_params": {"base_model": 42}}) == 42 + + # Adds the parent directory to the system path @@ -182,13 +189,13 @@ def test_potential_model_names_keeps_provider_prefixed_candidate(): Agent API serves `perplexity/glm-5.2`, mapped as `perplexity/perplexity/glm-5.2`) needs the un-stripped `/` candidate. Every other candidate reads the leading `perplexity/` as the litellm prefix and strips it away.""" - already_prefixed = _get_potential_model_names(model="perplexity/glm-5.2", custom_llm_provider="perplexity") + already_prefixed = get_potential_model_names(model="perplexity/glm-5.2", custom_llm_provider="perplexity") assert already_prefixed["provider_prefixed_model_name"] == "perplexity/perplexity/glm-5.2" assert already_prefixed["split_model"] == "glm-5.2" assert already_prefixed["combined_model_name"] == "perplexity/glm-5.2" assert already_prefixed["combined_stripped_model_name"] == "perplexity/glm-5.2" - bare = _get_potential_model_names(model="glm-5.2", custom_llm_provider="perplexity") + bare = get_potential_model_names(model="glm-5.2", custom_llm_provider="perplexity") assert bare["provider_prefixed_model_name"] == bare["combined_model_name"] == "perplexity/glm-5.2" @@ -242,9 +249,9 @@ def test_get_model_info_prefers_exact_dated_key_over_stripped( def test_get_model_info_internal_failure_is_not_reported_as_unmapped() -> None: - with patch("litellm.utils._get_potential_model_names", side_effect=RuntimeError("malformed metadata")): + with patch("litellm.utils.get_potential_model_names", side_effect=RuntimeError("malformed metadata")): with pytest.raises(Exception, match="This model isn't mapped yet") as exc_info: - litellm.utils._get_model_info_helper(model="gpt-4o", custom_llm_provider="openai") + litellm.utils.get_model_info_helper(model="gpt-4o", custom_llm_provider="openai") assert not isinstance(exc_info.value, litellm.ModelNotMappedError) @@ -4236,7 +4243,7 @@ async def test_s3_v2_success_callback_registers_alongside_user_subclass( and success_callback ["s3_v2"], the built-in s3_v2 logger was never added and S3 logs were silently dropped while requests kept returning 200.""" from litellm.integrations.s3_v2 import S3Logger - from litellm.utils import _add_custom_logger_callback_to_specific_event + from litellm.utils import add_custom_logger_callback_to_specific_event class UserS3Logger(S3Logger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -4248,7 +4255,7 @@ async def test_s3_v2_success_callback_registers_alongside_user_subclass( monkeypatch.setattr(litellm, "failure_callback", []) monkeypatch.setattr(litellm, "_async_failure_callback", []) - _add_custom_logger_callback_to_specific_event("s3_v2", "success") + add_custom_logger_callback_to_specific_event("s3_v2", "success") assert any(type(cb) is S3Logger for cb in litellm.success_callback) assert any(type(cb) is S3Logger for cb in litellm._async_success_callback) @@ -6012,16 +6019,16 @@ class TestDefaultReasoningEffortHydration: [("gpt-5.1", "openai"), ("gpt-5.4", "openai"), ("azure/gpt-5.1", "azure")], ) def test_the_declared_default_survives_model_info_hydration(self, local_model_cost_map, model, provider): - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper - model_info = dict(_get_model_info_helper(model=model, custom_llm_provider=provider)) + model_info = dict(get_model_info_helper(model=model, custom_llm_provider=provider)) assert model_info["default_reasoning_effort"] == "none" def test_a_model_that_declares_nothing_hydrates_to_none(self, local_model_cost_map): """Absent means "the map does not say", which the gate reads as reasoning being active.""" - from litellm.utils import _get_model_info_helper + from litellm.utils import get_model_info_helper - model_info = dict(_get_model_info_helper(model="gpt-5.6-terra", custom_llm_provider="openai")) + model_info = dict(get_model_info_helper(model="gpt-5.6-terra", custom_llm_provider="openai")) assert model_info.get("default_reasoning_effort") is None diff --git a/tests/unit/types/test_completion.py b/tests/unit/types/test_completion.py index 60928d3850b..28d13b7da4f 100644 --- a/tests/unit/types/test_completion.py +++ b/tests/unit/types/test_completion.py @@ -16,7 +16,7 @@ import pytest from litellm.types.completion import ( ChatCompletionMessageParam, CompletionRequest, - _CompletionDispatchContext, + CompletionDispatchContext, ) @@ -155,8 +155,8 @@ def test_completion_request_with_all_params(): assert request.n == 1 -def _build_dispatch_context() -> _CompletionDispatchContext: - return _CompletionDispatchContext( +def _build_dispatch_context() -> CompletionDispatchContext: + return CompletionDispatchContext( _azure_detection_model="gpt-4o", acompletion=False, api_base=None, diff --git a/tests/unit/types/test_prometheus_label_value_sanitize.py b/tests/unit/types/test_prometheus_label_value_sanitize.py index 9ff7eb460e0..234cabc540f 100644 --- a/tests/unit/types/test_prometheus_label_value_sanitize.py +++ b/tests/unit/types/test_prometheus_label_value_sanitize.py @@ -1,7 +1,7 @@ import pytest from litellm.types.integrations.prometheus import ( - _sanitize_prometheus_label_value, + sanitize_prometheus_label_value, ) @@ -30,5 +30,5 @@ from litellm.types.integrations.prometheus import ( ], ) def test_sanitize_prometheus_label_value_expected_outputs(value, expected): - assert _sanitize_prometheus_label_value(value) == expected + assert sanitize_prometheus_label_value(value) == expected diff --git a/tests/unit/vector_stores/test_vector_store_registry.py b/tests/unit/vector_stores/test_vector_store_registry.py index acfaccf8e2d..a29ce1463ab 100644 --- a/tests/unit/vector_stores/test_vector_store_registry.py +++ b/tests/unit/vector_stores/test_vector_store_registry.py @@ -147,7 +147,7 @@ def test_search_uses_registry_credentials(): litellm.vector_store_registry = registry try: logger = MagicMock() - logger._response_cost_calculator.return_value = 0 + logger.response_cost_calculator.return_value = 0 # Mock the search response mock_search_response = { @@ -300,7 +300,7 @@ def _database_listing_index_rows(rows: Sequence[object]) -> SimpleNamespace: ], ) async def test_vector_store_index_rows_from_the_db_are_returned_as_indexes(row: object) -> None: - indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db( + indexes: Final = await VectorStoreIndexRegistry.get_vector_store_indexes_from_db( _database_listing_index_rows([row]) ) @@ -309,7 +309,7 @@ async def test_vector_store_index_rows_from_the_db_are_returned_as_indexes(row: @pytest.mark.asyncio async def test_vector_store_index_row_without_optional_columns_gets_empty_defaults() -> None: - indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db( + indexes: Final = await VectorStoreIndexRegistry.get_vector_store_indexes_from_db( _database_listing_index_rows([{"id": "idx-1", "index_name": "team-docs", "litellm_params": _INDEX_PARAMS}]) ) @@ -342,4 +342,4 @@ async def test_vector_store_index_row_without_optional_columns_gets_empty_defaul ) async def test_malformed_vector_store_index_row_from_the_db_raises_a_validation_error(row: object) -> None: with pytest.raises(ValidationError): - await VectorStoreIndexRegistry._get_vector_store_indexes_from_db(_database_listing_index_rows([row])) + await VectorStoreIndexRegistry.get_vector_store_indexes_from_db(_database_listing_index_rows([row])) diff --git a/tests/vector_store_tests/base_vector_store_test.py b/tests/vector_store_tests/base_vector_store_test.py index 926fe98b6ec..1e61613decf 100644 --- a/tests/vector_store_tests/base_vector_store_test.py +++ b/tests/vector_store_tests/base_vector_store_test.py @@ -31,7 +31,7 @@ class BaseVectorStoreTest(ABC): @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_basic_search_vector_store(self, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True base_request_args = self.get_base_request_args() default_query = base_request_args.pop("query", "Basic ping") @@ -55,7 +55,7 @@ class BaseVectorStoreTest(ABC): @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_basic_create_vector_store(self, sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True base_request_args = self.get_base_create_vector_store_args() diff --git a/tests/vector_store_tests/rag/base_rag_tests.py b/tests/vector_store_tests/rag/base_rag_tests.py index caeb7651085..70a78144b80 100644 --- a/tests/vector_store_tests/rag/base_rag_tests.py +++ b/tests/vector_store_tests/rag/base_rag_tests.py @@ -77,7 +77,7 @@ class BaseRAGTest(ABC): """ Test basic text file ingestion to vector store. """ - litellm._turn_on_debug() + litellm.turn_on_debug() filename, unique_id = self.get_unique_filename("basic_ingest") text_content = f"Test document {unique_id} for RAG ingestion.".encode("utf-8") @@ -114,7 +114,7 @@ class BaseRAGTest(ABC): """ import asyncio - litellm._turn_on_debug() + litellm.turn_on_debug() filename, unique_id = self.get_unique_filename("ingest_query") text_content = f""" diff --git a/tests/vector_store_tests/rag/test_rag_openai.py b/tests/vector_store_tests/rag/test_rag_openai.py index 368e4e471b1..f90519a17e6 100644 --- a/tests/vector_store_tests/rag/test_rag_openai.py +++ b/tests/vector_store_tests/rag/test_rag_openai.py @@ -44,7 +44,7 @@ class TestRAGOpenAI(BaseRAGTest): """Test basic RAG query flow.""" import asyncio - litellm._turn_on_debug() + litellm.turn_on_debug() # First ingest a document filename, unique_id = self.get_unique_filename("rag_query") @@ -91,7 +91,7 @@ class TestRAGOpenAI(BaseRAGTest): """Test RAG query with reranking.""" import asyncio - litellm._turn_on_debug() + litellm.turn_on_debug() # First ingest a document filename, unique_id = self.get_unique_filename("rag_query_rerank") diff --git a/tests/vector_store_tests/rag/test_rag_vertex_ai.py b/tests/vector_store_tests/rag/test_rag_vertex_ai.py index ae5891ed3ff..67e0a25782b 100644 --- a/tests/vector_store_tests/rag/test_rag_vertex_ai.py +++ b/tests/vector_store_tests/rag/test_rag_vertex_ai.py @@ -142,7 +142,7 @@ class TestRAGVertexAI(BaseRAGTest): - Long-running operation polling for corpus creation - File upload to the newly created corpus """ - litellm._turn_on_debug() + litellm.turn_on_debug() filename, unique_id = self.get_unique_filename("create_corpus") text_content = f""" @@ -207,7 +207,7 @@ class TestRAGVertexAI(BaseRAGTest): if not corpus_id: pytest.skip("Skipping test: VERTEX_CORPUS_ID not set") - litellm._turn_on_debug() + litellm.turn_on_debug() filename, unique_id = self.get_unique_filename("existing_corpus") text_content = f""" diff --git a/tests/vector_store_tests/test_azure_ai_vector_store.py b/tests/vector_store_tests/test_azure_ai_vector_store.py index d1fc8436fc9..daa0c9b6ed0 100644 --- a/tests/vector_store_tests/test_azure_ai_vector_store.py +++ b/tests/vector_store_tests/test_azure_ai_vector_store.py @@ -20,7 +20,7 @@ from litellm.vector_stores import ( @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_basic_search_vector_store(sync_mode): - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True base_request_args = { "vector_store_id": "my-vector-index", diff --git a/tests/vector_store_tests/test_milvus_vector_store.py b/tests/vector_store_tests/test_milvus_vector_store.py index 2ba9168b49f..2139abb6a90 100644 --- a/tests/vector_store_tests/test_milvus_vector_store.py +++ b/tests/vector_store_tests/test_milvus_vector_store.py @@ -427,7 +427,7 @@ class TestMilvusVectorStore: # @pytest.mark.asyncio # async def test_basic_search_vector_store(sync_mode): # """Integration test with real Milvus API (requires credentials)""" -# litellm._turn_on_debug() +# litellm.turn_on_debug() # litellm.set_verbose = True # base_request_args = { # "vector_store_id": "book_2", diff --git a/tests/vector_store_tests/test_ragflow_vector_store.py b/tests/vector_store_tests/test_ragflow_vector_store.py index 0af821da98a..c0511ebacc6 100644 --- a/tests/vector_store_tests/test_ragflow_vector_store.py +++ b/tests/vector_store_tests/test_ragflow_vector_store.py @@ -316,7 +316,7 @@ class TestRAGFlowVectorStore(BaseVectorStoreTest): @pytest.mark.asyncio async def test_basic_create_vector_store(self, sync_mode): """Override to handle RAGFlow-specific connection errors.""" - litellm._turn_on_debug() + litellm.turn_on_debug() litellm.set_verbose = True base_request_args = self.get_base_create_vector_store_args()