From 58258409c912f00a07a7b0c43c6ce19bc7388e24 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 7 Oct 2026 18:59:27 -0700 Subject: [PATCH] refactor(proxy): expose public names for private proxy helpers (#45170) * refactor(proxy): expose public names for private proxy helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep original class names behind public aliases Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): preserve internal callback filtering Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep _PROXY_ class names for managed files hooks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep old private names in package exports Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): preserve recursive auth helper name Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): align MCP limiter tests with server enforcement Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): match main's MCP limiter tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep old private names bound in importing modules Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): preserve compatibility imports through strict lint Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): use exact pyright suppression in password helper test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(proxy): add reasons to compatibility import noqa comments Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): update IN-list baseline for renamed helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/hooks/managed_files.py | 4 +- .../proxy/hooks/managed_vector_stores.py | 4 +- litellm/batches/batch_utils.py | 4 +- litellm/google_genai/streaming_iterator.py | 2 +- .../SlackAlerting/slack_alerting.py | 8 +- litellm/integrations/gcs_pubsub/pub_sub.py | 8 +- litellm/integrations/prometheus.py | 4 +- litellm/integrations/shadow_eval_logger.py | 8 +- .../custom_logger_registry.py | 14 +- litellm/litellm_core_utils/litellm_logging.py | 22 +- .../prompt_templates/common_utils.py | 8 +- .../chat/guardrail_translation/handler.py | 4 +- .../context_management/editors/compact.py | 8 +- .../messages/streaming_iterator.py | 2 +- .../mcp_server/auth/user_api_key_auth_mcp.py | 101 +- .../mcp_server/bridge_token_flow.py | 43 +- .../mcp_server/byok_oauth_endpoints.py | 5 +- .../proxy/_experimental/mcp_server/catalog.py | 2 +- litellm/proxy/_experimental/mcp_server/db.py | 22 +- .../mcp_server/discoverable_endpoints.py | 76 +- .../_experimental/mcp_server/mcp_context.py | 12 +- .../_experimental/mcp_server/mcp_debug.py | 2 +- .../mcp_server/mcp_server_manager.py | 228 +- .../mcp_server/oauth2_flow_backfill.py | 8 +- .../mcp_server/oauth2_token_cache.py | 5 +- .../_experimental/mcp_server/oauth_utils.py | 5 +- .../mcp_server/openapi_to_mcp_generator.py | 32 +- .../_experimental/mcp_server/operations.py | 80 +- .../mcp_server/rest_endpoints.py | 61 +- .../mcp_server/sampling_handler.py | 4 +- .../mcp_server/semantic_tool_filter.py | 16 +- .../proxy/_experimental/mcp_server/server.py | 59 +- .../mcp_server/server_resolution.py | 16 +- .../mcp_server/ui_session_utils.py | 4 +- .../auth/agent_permission_handler.py | 4 +- litellm/proxy/agent_endpoints/endpoints.py | 4 +- .../claude_code_marketplace.py | 21 +- .../proxy/anthropic_endpoints/endpoints.py | 9 +- .../anthropic_endpoints/gateway_endpoints.py | 7 +- .../anthropic_endpoints/skills_endpoints.py | 8 +- litellm/proxy/auth/auth_checks.py | 227 +- .../proxy/auth/auth_checks_organization.py | 5 +- litellm/proxy/auth/auth_exception_handler.py | 11 +- litellm/proxy/auth/auth_utils.py | 15 +- litellm/proxy/auth/authorization.py | 4 +- litellm/proxy/auth/fallback_budget.py | 7 +- litellm/proxy/auth/ip_address_utils.py | 7 +- litellm/proxy/auth/resolvers/store.py | 23 +- litellm/proxy/auth/route_checks.py | 31 +- litellm/proxy/auth/user_api_key_auth.py | 191 +- litellm/proxy/batches_endpoints/endpoints.py | 22 +- litellm/proxy/common_request_processing.py | 69 +- .../auth_cache_invalidation_pubsub.py | 16 +- .../proxy/common_utils/cache_aware_routing.py | 4 +- litellm/proxy/common_utils/callback_utils.py | 16 +- .../proxy/common_utils/config_sync_pubsub.py | 20 +- .../common_utils/encrypt_decrypt_utils.py | 41 +- .../proxy/common_utils/http_parsing_utils.py | 36 +- .../common_utils/key_rotation_manager.py | 7 +- .../common_utils/openai_endpoint_utils.py | 7 +- .../common_utils/openapi_schema_compat.py | 4 +- litellm/proxy/common_utils/rbac_utils.py | 4 +- litellm/proxy/common_utils/realtime_utils.py | 6 +- .../proxy/container_endpoints/endpoints.py | 15 +- .../container_endpoints/handler_factory.py | 6 +- .../proxy/custom_hooks/custom_ui_sso_hook.py | 7 +- litellm/proxy/db/db_span.py | 9 +- litellm/proxy/db/db_spend_update_writer.py | 42 +- .../redis_update_buffer.py | 8 +- litellm/proxy/db/log_db_metrics.py | 7 +- .../proxy/decisions_endpoints/endpoints.py | 4 +- .../proxy/fine_tuning_endpoints/endpoints.py | 17 +- .../google_endpoints/agents_endpoints.py | 22 +- litellm/proxy/google_endpoints/endpoints.py | 27 +- .../proxy/guardrails/guardrail_endpoints.py | 11 +- .../guardrails/guardrail_hooks/azure/base.py | 6 +- .../guardrail_hooks/azure/prompt_shield.py | 9 +- .../guardrail_hooks/azure/text_moderation.py | 8 +- .../cisco_ai_defense/cisco_ai_defense.py | 7 +- .../cisco_ai_defense/cisco_ai_defense_mcp.py | 27 +- .../competitor_intent/airline.py | 31 +- .../competitor_intent/base.py | 19 +- .../mcp_end_user_permission.py | 2 +- .../guardrails/guardrail_hooks/presidio.py | 19 +- .../guardrails/guardrail_initializers.py | 4 +- .../proxy/guardrails/guardrail_registry.py | 7 +- litellm/proxy/health_check.py | 20 +- .../health_endpoints/_health_endpoints.py | 26 +- litellm/proxy/hooks/__init__.py | 42 +- litellm/proxy/hooks/azure_content_safety.py | 23 +- litellm/proxy/hooks/batch_enqueued_tokens.py | 4 +- litellm/proxy/hooks/batch_rate_limiter.py | 27 +- litellm/proxy/hooks/batch_redis_get.py | 9 +- litellm/proxy/hooks/cache_control_check.py | 5 +- litellm/proxy/hooks/dynamic_rate_limiter.py | 12 +- .../proxy/hooks/dynamic_rate_limiter_v3.py | 20 +- .../proxy/hooks/key_management_event_hooks.py | 7 +- .../hooks/max_budget_per_session_limiter.py | 5 +- litellm/proxy/hooks/max_iterations_limiter.py | 3 + .../proxy/hooks/mcp_semantic_filter/hook.py | 8 +- .../proxy/hooks/model_max_budget_limiter.py | 5 +- .../proxy/hooks/parallel_request_limiter.py | 13 +- .../hooks/parallel_request_limiter_v3.py | 27 +- .../proxy/hooks/prompt_injection_detection.py | 11 +- .../proxy/hooks/proxy_track_cost_callback.py | 22 +- litellm/proxy/hooks/sensitive_data_routing.py | 3 + litellm/proxy/image_endpoints/endpoints.py | 6 +- litellm/proxy/litellm_pre_call_utils.py | 79 +- .../access_group_endpoints.py | 32 +- .../auto_router_endpoints.py | 11 +- .../budget_management_endpoints.py | 9 +- .../cache_settings_endpoints.py | 18 +- .../management_endpoints/common_utils.py | 54 +- .../config_override_endpoints.py | 80 +- .../coordination_redis_endpoints.py | 10 +- .../credential_migration.py | 36 +- .../customer_endpoints.py | 7 +- .../internal_user_endpoints.py | 43 +- .../jwt_key_mapping_endpoints.py | 9 +- .../key_management_endpoints.py | 170 +- .../management_v1/spend_logs.py | 4 +- .../mcp_management_endpoints.py | 68 +- .../model_management_endpoints.py | 30 +- .../organization_endpoints.py | 35 +- .../policy_endpoints/__init__.py | 24 +- .../policy_endpoints/endpoints.py | 64 +- .../prompt_cache_prediction.py | 16 +- .../management_endpoints/scim/scim_v2.py | 21 +- .../management_endpoints/session_endpoints.py | 9 +- .../team_callback_endpoints.py | 25 +- .../management_endpoints/team_endpoints.py | 103 +- litellm/proxy/management_endpoints/ui_sso.py | 56 +- .../usage_endpoints/endpoints.py | 4 +- .../access_group_key_sync.py | 7 +- .../access_group_team_sync.py | 7 +- .../auto_router_permissions.py | 7 +- .../bulk_team_member_budgets.py | 7 +- .../management_helpers/bulk_user_creation.py | 14 +- .../management_helpers/bulk_user_deletion.py | 9 +- .../object_permission_utils.py | 29 +- .../team_member_permission_checks.py | 12 +- litellm/proxy/management_helpers/utils.py | 9 +- litellm/proxy/ocr_endpoints/endpoints.py | 2 +- .../proxy/openai_evals_endpoints/endpoints.py | 22 +- .../openai_files_endpoints/common_utils.py | 21 +- .../openai_files_endpoints/files_endpoints.py | 18 +- .../llm_passthrough_endpoints.py | 73 +- .../anthropic_passthrough_logging_handler.py | 16 +- .../assembly_passthrough_logging_handler.py | 16 +- .../openai_passthrough_logging_handler.py | 9 +- .../vertex_passthrough_logging_handler.py | 4 +- .../pass_through_endpoints.py | 35 +- .../passthrough_guardrails.py | 4 +- .../streaming_handler.py | 14 +- .../pass_through_endpoints/success_handler.py | 14 +- .../proxy/policy_engine/policy_registry.py | 16 +- .../proxy/policy_engine/policy_validator.py | 4 +- litellm/proxy/proxy_cli.py | 84 +- litellm/proxy/proxy_server.py | 497 +-- .../public_endpoints/public_endpoints.py | 16 +- .../public_endpoints/public_v1/model_hub.py | 4 +- litellm/proxy/rag_endpoints/endpoints.py | 14 +- litellm/proxy/realtime_endpoints/endpoints.py | 9 +- .../proxy/response_api_endpoints/endpoints.py | 68 +- litellm/proxy/route_llm_request.py | 4 +- litellm/proxy/search_endpoints/endpoints.py | 2 +- .../search_endpoints/search_tool_registry.py | 7 +- .../spend_tracking/budget_reservation.py | 28 +- .../spend_tracking/cloudzero_endpoints.py | 7 +- .../spend_management_endpoints.py | 30 +- .../spend_tracking/spend_tracking_utils.py | 10 +- .../proxy/spend_tracking/vantage_endpoints.py | 7 +- .../proxy_setting_endpoints.py | 8 +- litellm/proxy/utils.py | 202 +- .../proxy/vector_store_endpoints/endpoints.py | 24 +- .../vector_store_files_endpoints/endpoints.py | 20 +- .../vertex_ai_endpoints/langfuse_endpoints.py | 14 +- litellm/proxy/video_endpoints/endpoints.py | 33 +- .../mcp/litellm_proxy_mcp_handler.py | 6 +- .../responses/mcp/mcp_streaming_iterator.py | 8 +- litellm/responses/utils.py | 14 +- tests/batches_tests/test_batch_rate_limits.py | 24 +- .../unbounded_in_baseline.txt | 4 +- tests/guardrails_tests/test_presidio_pii.py | 10 +- ...st_redis_ttl_preserving_token_increment.py | 2 +- tests/local_testing/test_caching.py | 4 +- .../test_pass_through_unit_tests.py | 2 +- .../cache/test_python_cache.py | 12 +- tests/test_presidio_latency.py | 14 +- tests/unit/batches/test_batch_utils.py | 4 +- .../test_request_redis_batch_post_call.py | 6 +- .../test_request_redis_batch_pre_call.py | 12 +- .../test_container_proxy_ownership.py | 12 +- .../test_prometheus_logging_callbacks.py | 4 +- .../proxy/auth/test_user_api_key_auth.py | 16 +- .../proxy/hooks/test_managed_files.py | 12 +- .../test_batch_retrieve_input_file_id.py | 4 +- .../test_otel_admin_endpoints.py | 34 +- .../integrations/test_shadow_eval_logger.py | 2 +- .../test_health_check_helpers.py | 42 +- .../test_litellm_logging.py | 37 +- .../test_anthropic_guardrail_handler.py | 8 +- .../context_management/test_compact.py | 14 +- .../messages/test_response_cache.py | 2 +- .../auth/test_user_api_key_auth_mcp.py | 92 +- .../mcp_server/test_db_credentials.py | 4 +- .../mcp_server/test_discoverable_endpoints.py | 138 +- .../mcp_server/test_jwt_mcp_enforcement.py | 4 +- .../mcp_server/test_jwt_mcp_simple.py | 4 +- .../mcp_server/test_mcp_env_vars.py | 2 +- .../test_mcp_guardrail_usage_monitor.py | 4 +- .../mcp_server/test_mcp_hook_extra_headers.py | 102 +- .../test_mcp_max_concurrent_requests.py | 2 +- .../test_mcp_oauth_passthrough_tools.py | 14 +- .../mcp_server/test_mcp_proxy_mode.py | 2 +- .../mcp_server/test_mcp_server.py | 48 +- .../mcp_server/test_mcp_server_manager.py | 286 +- .../test_mcp_server_tool_calls_and_headers.py | 146 +- .../mcp_server/test_mcp_sigv4_auth.py | 20 +- .../mcp_server/test_mcp_tool_search.py | 2 +- .../mcp_server/test_mcp_toolset_scope.py | 54 +- .../mcp_server/test_oauth2_token_cache.py | 22 +- .../test_openapi_to_mcp_generator.py | 70 +- .../mcp_server/test_openapi_tool_auth.py | 52 +- .../mcp_server/test_operations.py | 16 +- .../mcp_server/test_per_user_oauth_cache.py | 10 +- .../mcp_server/test_rest_endpoints.py | 66 +- .../mcp_server/test_server_resolution.py | 14 +- .../auth/test_agent_permission_handler.py | 6 +- .../anthropic_endpoints/test_endpoints.py | 14 +- tests/unit/proxy/auth/test_auth_checks.py | 72 +- ...st_auth_checks_object_access_and_lookup.py | 276 +- .../proxy/auth/test_auth_exception_handler.py | 34 +- tests/unit/proxy/auth/test_auth_utils.py | 6 +- .../auth/test_custom_auth_end_user_budget.py | 4 +- .../test_default_end_user_budget_simple.py | 4 +- tests/unit/proxy/auth/test_login_utils.py | 4 +- .../auth/test_object_permission_loading.py | 6 +- tests/unit/proxy/auth/test_proxy_routes.py | 4 +- tests/unit/proxy/auth/test_route_checks.py | 26 +- .../test_router_override_fallback_auth.py | 14 +- .../test_unmapped_model_budget_enforcement.py | 78 +- .../unit/proxy/auth/test_user_api_key_auth.py | 12 +- .../test_user_api_key_auth_request_flow.py | 262 +- .../proxy/batches_endpoints/test_endpoints.py | 38 +- .../test_litellm_executed_batches.py | 4 +- .../common_utils/test_cache_aware_routing.py | 4 +- .../common_utils/test_check_batch_cost.py | 13 +- .../test_encrypt_decrypt_utils.py | 32 +- .../common_utils/test_http_parsing_utils.py | 106 +- .../proxy/common_utils/test_rbac_utils.py | 16 +- .../proxy/common_utils/test_realtime_cache.py | 18 +- .../test_upsert_budget_membership.py | 32 +- .../test_redis_update_buffer.py | 4 +- .../proxy/db/test_db_spend_update_writer.py | 6 +- .../proxy/db/test_update_daily_tag_spend.py | 28 +- .../test_gemini_agents_endpoints.py | 90 +- .../guardrail_hooks/noma/test_noma.py | 4 +- .../test_bedrock_guardrails.py | 2 +- .../guardrail_hooks/test_presidio.py | 222 +- .../test_deferred_guardrail_logging.py | 6 +- .../guardrails/test_guardrail_coverage.py | 8 +- .../proxy/guardrails/test_init_guardrails.py | 4 +- .../health_endpoints/test_health_endpoints.py | 30 +- .../proxy/hooks/test_batch_file_validation.py | 178 +- .../proxy/hooks/test_batch_rate_limiter.py | 6 +- .../proxy/hooks/test_dynamic_rate_limiter.py | 8 +- .../hooks/test_dynamic_rate_limiter_v3.py | 2 +- .../test_max_budget_per_session_limiter.py | 12 +- .../hooks/test_max_iterations_limiter.py | 16 +- .../hooks/test_model_max_budget_limiter.py | 18 +- .../hooks/test_parallel_request_limiter.py | 6 +- .../hooks/test_parallel_request_limiter_v3.py | 31 +- .../hooks/test_prompt_injection_detection.py | 20 +- .../unit/proxy/hooks/test_proxy_hooks_init.py | 4 +- .../test_proxy_rate_limit_provider_field.py | 149 +- .../hooks/test_proxy_track_cost_callback.py | 100 +- .../proxy/hooks/test_rate_limiter_toctou.py | 10 +- .../hooks/test_sensitive_data_routing.py | 32 +- tests/unit/proxy/hooks/test_tpm_concurrent.py | 2 +- ...test_unit_test_max_model_budget_limiter.py | 64 +- .../scim/test_scim_key_deactivation.py | 20 +- .../scim/test_scim_v2_endpoints.py | 4 +- .../test_auto_router_endpoints.py | 2 +- .../test_cache_settings_endpoints.py | 50 +- .../management_endpoints/test_common_utils.py | 196 +- .../test_config_override_endpoints.py | 16 +- .../test_coordination_redis_endpoints.py | 10 +- .../test_credential_migration.py | 12 +- .../test_delete_verification_tokens_failed.py | 8 +- .../test_internal_user_endpoints.py | 6 +- .../test_key_generate_prisma.py | 52 +- .../test_key_management_endpoints.py | 184 +- .../test_mcp_management_endpoints.py | 100 +- .../test_model_management_endpoints.py | 68 +- .../test_org_admin_team_access.py | 46 +- .../test_organization_endpoints.py | 16 +- .../test_policy_endpoints.py | 98 +- .../test_prompt_cache_prediction.py | 12 +- .../test_ptu_model_settings.py | 10 +- .../test_session_endpoints.py | 10 +- .../test_team_callback_endpoints.py | 28 +- .../test_team_endpoints.py | 58 +- .../test_team_model_alias_merge.py | 14 +- .../proxy/management_endpoints/test_ui_sso.py | 12 +- .../test_management_helpers_utils.py | 4 +- .../test_object_permission_utils.py | 120 +- .../test_team_metadata_validation.py | 2 +- .../test_batch_guardrails.py | 4 +- ...t_anthropic_passthrough_logging_handler.py | 20 +- ...test_gemini_passthrough_logging_handler.py | 6 +- ...test_openai_passthrough_logging_handler.py | 4 +- ...st_typesafe_passthrough_logging_handler.py | 4 +- .../test_deepgram_ws_passthrough_routes.py | 4 +- .../test_llm_pass_through_endpoints.py | 38 +- .../test_pass_through_endpoints.py | 24 +- ...t_passthrough_guardrail_block_otel_span.py | 4 +- .../test_passthrough_post_call_guardrails.py | 4 +- .../test_streaming_handler_interrupt.py | 18 +- .../proxy_server/test_background_health.py | 8 +- .../unit/proxy/proxy_server/test_lifecycle.py | 18 +- .../proxy/proxy_server/test_proxy_config.py | 6 +- .../test_routes_chat_completions.py | 14 +- .../proxy/proxy_server/test_routes_config.py | 4 +- .../proxy_server/test_routes_embeddings.py | 14 +- .../proxy_server/test_routes_invitation.py | 8 +- .../proxy_server/test_routes_model_info.py | 2 +- .../proxy_server/test_streaming_helpers.py | 2 +- .../public_endpoints/test_public_endpoints.py | 14 +- .../response_api_endpoints/test_endpoints.py | 30 +- .../spend_tracking/test_search_api_logging.py | 4 +- .../test_spend_management_endpoints.py | 4 +- .../test_spend_tracking_utils.py | 40 +- .../proxy/test_aiohttp_session_recovery.py | 8 +- .../test_batch_x_litellm_model_encoding.py | 6 +- tests/unit/proxy/test_budget_reservation.py | 18 +- .../proxy/test_chat_completion_metadata.py | 54 +- .../proxy/test_common_request_processing.py | 44 +- tests/unit/proxy/test_dynamic_mcp_route.py | 6 +- .../unit/proxy/test_health_check_functions.py | 14 +- .../proxy/test_health_check_max_tokens.py | 72 +- .../unit/proxy/test_litellm_pre_call_utils.py | 44 +- .../unit/proxy/test_model_level_guardrails.py | 18 +- .../proxy/test_model_list_callback_filter.py | 2 +- .../proxy/test_model_list_discoverable.py | 2 +- .../proxy/test_model_list_healthy_only.py | 4 +- ...t_modify_response_streaming_passthrough.py | 2 +- tests/unit/proxy/test_proxy_cli.py | 48 +- tests/unit/proxy/test_proxy_server.py | 4 +- ...test_proxy_server_endpoints_and_startup.py | 68 +- .../proxy/test_proxy_setting_guardrails.py | 2 +- tests/unit/proxy/test_proxy_token_counter.py | 36 +- tests/unit/proxy/test_proxy_utils.py | 20 +- ..._utils_model_creation_and_error_logging.py | 20 +- .../test_response_polling_pre_call_checks.py | 16 +- tests/unit/proxy/test_team_member_update.py | 24 +- .../unit/proxy/test_unit_test_proxy_hooks.py | 2 +- tests/unit/proxy/test_update_spend.py | 2 +- .../test_zero_cost_model_budget_bypass.py | 18 +- .../test_proxy_setting_endpoints.py | 40 +- .../utils/helpers/test_guardrail_merge.py | 36 +- .../helpers/test_month_end_projection.py | 47 +- .../utils/helpers/test_premium_user_check.py | 10 +- .../proxy/utils/helpers/test_team_configs.py | 10 +- .../proxy/utils/helpers/test_url_helpers.py | 24 +- .../proxy/utils/prisma_and_spend/conftest.py | 2 +- .../prisma_and_spend/test_cache_user_row.py | 10 +- .../prisma_and_spend/test_password_helpers.py | 20 +- .../prisma_and_spend/test_spend_functions.py | 46 +- .../utils/proxy_logging/test_pre_call_hook.py | 8 +- .../test_vector_store_tenant_guard.py | 46 +- .../proxy/video_endpoints/test_endpoints.py | 4 +- .../mcp/test_litellm_proxy_mcp_handler.py | 18 +- .../mcp/test_mcp_streaming_iterator.py | 2 +- tests/unit/test_private_usage_aliases.py | 3029 ++++++++++++++++- .../unit/test_rate_limit_error_unification.py | 76 +- tests/unit/test_video_generation.py | 4 +- 377 files changed, 8886 insertions(+), 4979 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index c07bb216602..28c4cf5e66f 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -255,7 +255,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 @@ -2110,4 +2110,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 +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 c53f359a468..0e9657aba96 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py @@ -38,7 +38,7 @@ else: PrismaClient = Any -class PROXY_LiteLLMManagedVectorStores( +class _PROXY_LiteLLMManagedVectorStores( CustomLogger, BaseManagedResource[VectorStoreCreateResponse] ): """ @@ -462,4 +462,4 @@ class PROXY_LiteLLMManagedVectorStores( parent_otel_span=parent_otel_span, resource_id_key="vector_store_id", ) -_PROXY_LiteLLMManagedVectorStores = PROXY_LiteLLMManagedVectorStores +PROXY_LiteLLMManagedVectorStores = _PROXY_LiteLLMManagedVectorStores diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index ba625ac5c00..6573f13be75 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -416,11 +416,11 @@ def _provider_output_file_id(output_file_id: str) -> str: llm_output_file_id, model-encoded ids decode to the raw provider id, raw ids pass through. """ from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, get_original_file_id, + is_base64_encoded_unified_file_id, ) - unified_file_id: Final = _is_base64_encoded_unified_file_id(output_file_id) + unified_file_id: Final = is_base64_encoded_unified_file_id(output_file_id) if not unified_file_id: return get_original_file_id(output_file_id) try: diff --git a/litellm/google_genai/streaming_iterator.py b/litellm/google_genai/streaming_iterator.py index 27883c938db..64508ec93a6 100644 --- a/litellm/google_genai/streaming_iterator.py +++ b/litellm/google_genai/streaming_iterator.py @@ -90,7 +90,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: end_time: Final = datetime.now() asyncio.create_task( - PassThroughStreamingHandler._route_streaming_logging_to_handler( + PassThroughStreamingHandler.route_streaming_logging_to_handler( litellm_logging_obj=self.litellm_logging_obj, passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, url_route="/v1/generateContent", diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 704a25b0457..7e9a8fe877c 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -1845,7 +1845,7 @@ Model Info: try: from litellm.proxy.spend_tracking.spend_management_endpoints import ( - _get_spend_report_for_time_range, + get_spend_report_for_time_range, ) # Parse the time range @@ -1862,7 +1862,7 @@ Model Info: if await self.internal_usage_cache.async_get_cache(key=_event_cache_key): return - _resp: Final = await _get_spend_report_for_time_range( + _resp: Final = await get_spend_report_for_time_range( start_date=start_date.strftime("%Y-%m-%d"), end_date=todays_date.strftime("%Y-%m-%d"), ) @@ -1909,7 +1909,7 @@ Model Info: from calendar import monthrange from litellm.proxy.spend_tracking.spend_management_endpoints import ( - _get_spend_report_for_time_range, + get_spend_report_for_time_range, ) todays_date: Final = datetime.datetime.now().date() @@ -1921,7 +1921,7 @@ Model Info: if await self.internal_usage_cache.async_get_cache(key=_event_cache_key): return - _resp: Final = await _get_spend_report_for_time_range( + _resp: Final = await get_spend_report_for_time_range( start_date=first_day_of_month.strftime("%Y-%m-%d"), end_date=last_day_of_month.strftime("%Y-%m-%d"), ) diff --git a/litellm/integrations/gcs_pubsub/pub_sub.py b/litellm/integrations/gcs_pubsub/pub_sub.py index 293be174811..c799245c215 100644 --- a/litellm/integrations/gcs_pubsub/pub_sub.py +++ b/litellm/integrations/gcs_pubsub/pub_sub.py @@ -44,9 +44,9 @@ class GcsPubSubLogger(CustomBatchLogger): topic_id (str): Pub/Sub topic ID credentials_path (str, optional): Path to Google Cloud credentials JSON file """ - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check - _premium_user_check() + premium_user_check() self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) @@ -107,9 +107,9 @@ class GcsPubSubLogger(CustomBatchLogger): from litellm.proxy.spend_tracking.spend_tracking_utils import ( get_logging_payload, ) - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check - _premium_user_check() + premium_user_check() try: verbose_logger.debug("PubSub: Logging - Enters logging function for model %s", kwargs) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index ede75381277..e11ce2bf472 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -3833,7 +3833,7 @@ class PrometheusLogger(CustomLogger): """ from litellm.constants import UI_SESSION_TOKEN_TEAM_ID from litellm.proxy.management_endpoints.key_management_endpoints import ( - _list_key_helper, + list_key_helper, ) from litellm.proxy.proxy_server import prisma_client @@ -3847,7 +3847,7 @@ class PrometheusLogger(CustomLogger): list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken], int | None, ]: - key_list_response: Final = await _list_key_helper( + key_list_response: Final = await list_key_helper( prisma_client=prisma_client, page=page, size=page_size, diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index 76ba1d80841..046f510266a 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -633,9 +633,9 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool: from litellm.exceptions import BudgetExceededError from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import ( - _team_max_budget_check, - _virtual_key_max_budget_check, get_team_object, + team_max_budget_check, + virtual_key_max_budget_check, ) from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache except ImportError: @@ -645,7 +645,7 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool: if not isinstance(auth, UserAPIKeyAuth): return False try: - await _virtual_key_max_budget_check(valid_token=auth, proxy_logging_obj=proxy_logging_obj) + await virtual_key_max_budget_check(valid_token=auth, proxy_logging_obj=proxy_logging_obj) if auth.team_id: team: Final = await get_team_object( team_id=auth.team_id, @@ -653,7 +653,7 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool: user_api_key_cache=user_api_key_cache, check_cache_only=True, ) - await _team_max_budget_check(team_object=team, valid_token=auth, proxy_logging_obj=proxy_logging_obj) + await team_max_budget_check(team_object=team, valid_token=auth, proxy_logging_obj=proxy_logging_obj) except BudgetExceededError: return True except Exception as e: # noqa: BLE001 # advisory gate: a failed read must not block sampling diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index 1d277995211..8367c9ee00e 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -53,8 +53,14 @@ from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook i VectorStorePreCallHook, ) from litellm.integrations.zerobus import ZerobusLogger -from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler -from litellm.proxy.hooks.dynamic_rate_limiter_v3 import _PROXY_DynamicRateLimitHandlerV3 +from litellm.proxy.hooks.dynamic_rate_limiter import ( # noqa: F401 # legacy module exports + PROXY_DynamicRateLimitHandler, + _PROXY_DynamicRateLimitHandler, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) +from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( # noqa: F401 # legacy module exports + PROXY_DynamicRateLimitHandlerV3, + _PROXY_DynamicRateLimitHandlerV3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) class CustomLoggerRegistry: @@ -101,8 +107,8 @@ class CustomLoggerRegistry: "pointfive": PointFiveLogger, "zerobus": ZerobusLogger, "aws_sqs": SQSLogger, - "dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler, - "dynamic_rate_limiter_v3": _PROXY_DynamicRateLimitHandlerV3, + "dynamic_rate_limiter": PROXY_DynamicRateLimitHandler, + "dynamic_rate_limiter_v3": PROXY_DynamicRateLimitHandlerV3, "vector_store_pre_call_hook": VectorStorePreCallHook, "dotprompt": DotpromptManager, "bitbucket": BitBucketPromptManager, diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index d0b8109441a..eb11bd9c11c 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -4969,17 +4969,17 @@ def _init_custom_logger_compatible_class( return _otel_logger elif logging_integration == "dynamic_rate_limiter": from litellm.proxy.hooks.dynamic_rate_limiter import ( - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_DynamicRateLimitHandler): + if isinstance(callback, PROXY_DynamicRateLimitHandler): return callback if internal_usage_cache is None: raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}") - dynamic_rate_limiter_obj: Final = _PROXY_DynamicRateLimitHandler(internal_usage_cache=internal_usage_cache) + dynamic_rate_limiter_obj: Final = PROXY_DynamicRateLimitHandler(internal_usage_cache=internal_usage_cache) if llm_router is not None and isinstance(llm_router, litellm.Router): dynamic_rate_limiter_obj.update_variables(llm_router=llm_router) @@ -4987,17 +4987,19 @@ def _init_custom_logger_compatible_class( return dynamic_rate_limiter_obj elif logging_integration == "dynamic_rate_limiter_v3": from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3, + PROXY_DynamicRateLimitHandlerV3, ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_DynamicRateLimitHandlerV3): + if isinstance(callback, PROXY_DynamicRateLimitHandlerV3): return callback if internal_usage_cache is None: raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}") - dynamic_rate_limiter_obj_v3 = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=internal_usage_cache) + dynamic_rate_limiter_obj_v3: Final = PROXY_DynamicRateLimitHandlerV3( + internal_usage_cache=internal_usage_cache + ) if llm_router is not None and isinstance(llm_router, litellm.Router): dynamic_rate_limiter_obj_v3.update_variables(llm_router=llm_router) @@ -5546,19 +5548,19 @@ def get_custom_logger_compatible_class( elif logging_integration == "dynamic_rate_limiter": from litellm.proxy.hooks.dynamic_rate_limiter import ( - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_DynamicRateLimitHandler): + if isinstance(callback, PROXY_DynamicRateLimitHandler): return callback elif logging_integration == "dynamic_rate_limiter_v3": from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3, + PROXY_DynamicRateLimitHandlerV3, ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_DynamicRateLimitHandlerV3): + if isinstance(callback, PROXY_DynamicRateLimitHandlerV3): return callback elif logging_integration == "langtrace": diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 1f75125df09..0243141eda2 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -565,9 +565,9 @@ def update_messages_with_model_file_ids( } """ from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, convert_b64_uid_to_unified_uid, get_original_file_id, + is_base64_encoded_unified_file_id, is_model_embedded_id, ) @@ -603,7 +603,7 @@ def update_messages_with_model_file_ids( if model_file_id_mapping and model_id is not None else None ) - if not provider_file_id and _is_base64_encoded_unified_file_id(file_id): + if not provider_file_id and is_base64_encoded_unified_file_id(file_id): unified_file_id = convert_b64_uid_to_unified_uid(file_id) if "llm_output_file_id," in unified_file_id: provider_file_id = unified_file_id.split("llm_output_file_id,")[1].split(";")[0] @@ -634,9 +634,9 @@ def update_responses_input_with_model_file_ids( Format: {"litellm_file_id": {"model_id": "provider_file_id"}} """ from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, convert_b64_uid_to_unified_uid, get_original_file_id, + is_base64_encoded_unified_file_id, is_model_embedded_id, ) @@ -671,7 +671,7 @@ def update_responses_input_with_model_file_ids( updated_content.append(updated_content_item) else: # Check if this is a base64-encoded unified file ID without mapping - is_unified_file_id = _is_base64_encoded_unified_file_id(file_id) + is_unified_file_id = is_base64_encoded_unified_file_id(file_id) if is_unified_file_id: # Fallback: decode unified file ID unified_file_id = convert_b64_uid_to_unified_uid(file_id) diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 405ebce548a..e3442bca1e7 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -322,7 +322,7 @@ class AnthropicMessagesHandler(BaseTranslation): if not chunks: return None try: - return AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks( + return AnthropicPassthroughLoggingHandler.build_usage_only_response_from_chunks( all_chunks=chunks, model=str((request_data or {}).get("model") or ""), ) @@ -1226,7 +1226,7 @@ class AnthropicMessagesHandler(BaseTranslation): has_ended: Final = self._check_streaming_has_ended(responses_so_far) if has_ended: # build the model response from the responses_so_far - built_response: Final = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( + built_response: Final = AnthropicPassthroughLoggingHandler.build_complete_streaming_response( all_chunks=responses_so_far, litellm_logging_obj=cast("LiteLLMLoggingObj", litellm_logging_obj), model="", diff --git a/litellm/llms/anthropic/pass_through/context_management/editors/compact.py b/litellm/llms/anthropic/pass_through/context_management/editors/compact.py index 62826865894..f743320b702 100644 --- a/litellm/llms/anthropic/pass_through/context_management/editors/compact.py +++ b/litellm/llms/anthropic/pass_through/context_management/editors/compact.py @@ -228,7 +228,7 @@ async def _check_summary_model_access( try: from litellm.proxy._types import ProxyException from litellm.proxy.auth.auth_checks import ( - _can_object_call_model, + can_object_call_model, can_project_access_model, can_user_call_model, get_project_object, @@ -258,7 +258,7 @@ async def _check_summary_model_access( if not models: continue try: - _can_object_call_model( + can_object_call_model( model=summary_model, llm_router=llm_router, models=models, @@ -370,7 +370,7 @@ async def _check_summary_model_access( ) if member_allowed_models: try: - _can_object_call_model( + can_object_call_model( model=summary_model, llm_router=llm_router, models=list(member_allowed_models), @@ -558,7 +558,7 @@ async def _check_summary_model_rate_limit( limiter: Final[object] = get_proxy_hook("parallel_request_limiter") if get_proxy_hook is not None else None should_rate_limit_check: Final[_ShouldRateLimit | None] = getattr(limiter, "should_rate_limit", None) create_descriptors: Final[_CreateRateLimitDescriptors | None] = getattr( - limiter, "_create_rate_limit_descriptors", None + limiter, "create_rate_limit_descriptors", None ) add_team_descriptor: Final[_AddModelRateLimitDescriptor | None] = getattr( limiter, "_add_team_model_rate_limit_descriptor_from_metadata", None diff --git a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py index 25e0f93b04c..f37e5190850 100644 --- a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py @@ -421,7 +421,7 @@ class BaseAnthropicMessagesStreamingIterator: if self.completion_start_time is not None: self.litellm_logging_obj.completion_start_time = self.completion_start_time self.litellm_logging_obj.model_call_details["completion_start_time"] = self.completion_start_time - logging_coroutine: Final = PassThroughStreamingHandler._route_streaming_logging_to_handler( + logging_coroutine: Final = PassThroughStreamingHandler.route_streaming_logging_to_handler( litellm_logging_obj=self.litellm_logging_obj, passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, url_route="/v1/messages", diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 418bdcd2941..ab14a317a66 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -56,12 +56,16 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import ( from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy from litellm.proxy.auth.ip_address_utils import IPAddressUtils -from litellm.proxy.auth.user_api_key_auth import ( +from litellm.proxy.auth.user_api_key_auth import ( # noqa: F401 # legacy module exports _get_bearer_token_or_received_api_key, # pyright: ignore[reportPrivateUsage] # shared x-litellm-api-key parser lives with user_api_key_auth - _run_centralized_common_checks, + _run_centralized_common_checks, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + run_centralized_common_checks, user_api_key_auth, ) -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.user_api_key_cache import ( AUTH_OBJECTS_TARGET, USER_NO_MCP_PERMISSION_SENTINEL, @@ -225,7 +229,7 @@ def _agent_capped_servers( ) -def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool: +def is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool: """True when this auth is a keyless subject admitted by the gateway session / bridge user path, as opposed to a JWT or other keyless auth that merely lacks a ``team_id``. @@ -235,6 +239,9 @@ def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> b return user_api_key_auth is not None and user_api_key_auth.mcp_admitted_user_subject is True +_is_mcp_admitted_user_subject: Final = is_mcp_admitted_user_subject + + def _gateway_dcr_challenge_target( route: str, mcp_servers: list[str] | None, @@ -452,7 +459,7 @@ class MCPRequestHandler: HTTPException: If headers are invalid or missing required headers """ async with global_manager().catalog.operation(): - headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope) + headers: Final = MCPRequestHandler.safe_get_headers_from_scope(scope) # Check if there is an explicit LiteLLM API key (primary header) has_explicit_litellm_key: Final = ( @@ -462,13 +469,19 @@ class MCPRequestHandler: litellm_api_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(headers) or "" # Get the old mcp_auth_header for backward compatibility - mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers) + mcp_auth_header = MCPRequestHandler.get_mcp_auth_header_from_headers( # rebind-ok: pre-existing rebinding on a rename-only line + headers + ) # Get the new server-specific auth headers - mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers) + mcp_server_auth_headers = MCPRequestHandler.get_mcp_server_auth_headers_from_headers( # rebind-ok: pre-existing rebinding on a rename-only line + headers + ) # Get the oauth2 headers - oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers) + oauth2_headers = MCPRequestHandler.get_oauth2_headers_from_headers( # rebind-ok: pre-existing rebinding on a rename-only line + headers + ) # Parse MCP servers from header mcp_servers_header: Final = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME) @@ -621,7 +634,7 @@ class MCPRequestHandler: mcp_auth_header, mcp_server_auth_headers, ) = MCPRequestHandler._scrub_gateway_admission_credentials( - admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth), + admitted=is_mcp_admitted_user_subject(validated_user_api_key_auth), oauth2_headers=oauth2_headers, raw_headers=raw_headers, mcp_auth_header=mcp_auth_header, @@ -1060,7 +1073,7 @@ class MCPRequestHandler: await pre_db_read_auth_checks( request=request, - request_data=await _read_request_body(request=request), + request_data=await read_request_body(request=request), route=route, ) @@ -1351,10 +1364,10 @@ class MCPRequestHandler: admitted.budget_reservation = None try: RouteChecks.should_call_route(route=route, valid_token=admitted, request=request) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=admitted, request=request, - request_data=await _read_request_body(request=request), + request_data=await read_request_body(request=request), route=route, ) except (HTTPException, ProxyException): @@ -1401,7 +1414,7 @@ class MCPRequestHandler: return mcp_servers_header if mcp_servers_header is not None else [] @staticmethod - def _get_mcp_auth_header_from_headers(headers: Headers) -> str | None: + def get_mcp_auth_header_from_headers(headers: Headers) -> str | None: """ Get the header passed to LiteLLM to pass to downstream MCP servers @@ -1424,8 +1437,10 @@ class MCPRequestHandler: ) return auth_header + _get_mcp_auth_header_from_headers = get_mcp_auth_header_from_headers + @staticmethod - def _get_mcp_server_auth_headers_from_headers( + def get_mcp_server_auth_headers_from_headers( headers: Headers, ) -> dict[str, dict[str, str]]: """ @@ -1478,8 +1493,10 @@ class MCPRequestHandler: return server_auth_headers + _get_mcp_server_auth_headers_from_headers = get_mcp_server_auth_headers_from_headers + @staticmethod - def _get_oauth2_headers_from_headers(headers: Headers) -> dict[str, str]: + def get_oauth2_headers_from_headers(headers: Headers) -> dict[str, str]: """ Get the oauth2 headers from the request headers. """ @@ -1489,6 +1506,8 @@ class MCPRequestHandler: oauth2_headers["Authorization"] = header_value return oauth2_headers + _get_oauth2_headers_from_headers = get_oauth2_headers_from_headers + @staticmethod def get_mcp_client_side_auth_header_name() -> str: """ @@ -1535,7 +1554,7 @@ class MCPRequestHandler: return None @staticmethod - def _safe_get_headers_from_scope(scope: Scope) -> Headers: + def safe_get_headers_from_scope(scope: Scope) -> Headers: """ Safely extract headers from ASGI scope using Starlette's Headers class which handles case insensitivity and proper header parsing. @@ -1563,6 +1582,8 @@ class MCPRequestHandler: # Return empty Headers object with empty dict return Headers({}) + _safe_get_headers_from_scope = safe_get_headers_from_scope + @staticmethod def _reject_duplicate_authorization(raw_headers: object) -> None: """Raise 400 when the raw ASGI headers carry more than one ``Authorization`` header.""" @@ -1638,7 +1659,7 @@ class MCPRequestHandler: # matters: the no_mcp_servers opt-out below reads the caller's own object_permission, so above # this branch a user's own opt-out would wrongly zero their TEAMS' grants too (each source is # independent; an opt-out silences only its own source, inside the recursive call). - if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None: + if is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None: return MCPServerAccess( server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)), ) @@ -1950,18 +1971,18 @@ class MCPRequestHandler: # already-exceeded state; ATTRIBUTION of new spend stays with the user (documented deferral). from litellm.exceptions import BudgetExceededError from litellm.proxy.auth.auth_checks import ( - _organization_max_budget_check, - _team_max_budget_check, + organization_max_budget_check, + team_max_budget_check, ) source_view: Final = MCPRequestHandler._scoped_source_auth( auth, team_id=team_id, org_id=team_obj.organization_id or auth.org_id, carry_user_grants=False ) try: - await _team_max_budget_check( + await team_max_budget_check( team_object=team_obj, valid_token=source_view, proxy_logging_obj=proxy_logging_obj ) - await _organization_max_budget_check( + await organization_max_budget_check( valid_token=source_view, team_object=team_obj, prisma_client=prisma_client, @@ -2042,14 +2063,14 @@ class MCPRequestHandler: owning ``org_id`` so the team's budget accumulates and the right org is charged. Falls back to user-level attribution (rather than guessing a team) when the tool name does not resolve to a server, reusing the manager's own tool-name lookup.""" - if not _is_mcp_admitted_user_subject(auth): + if not is_mcp_admitted_user_subject(auth): return auth try: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) - server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name) + server: Final = global_mcp_server_manager.get_mcp_server_from_tool_name(tool_name) if server is None: return auth source: Final = await MCPRequestHandler.attributing_source_for_server(auth, server.server_id) @@ -2330,7 +2351,7 @@ class MCPRequestHandler: # FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per # source and shares nothing with the single-credential prelude below. Ordering is the invariant: # sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant. - if _is_mcp_admitted_user_subject(user_api_key_auth): + if is_mcp_admitted_user_subject(user_api_key_auth): return await MCPRequestHandler.resolve_admitted_subject_tools(server_id, user_api_key_auth) # Get key and team object permissions (already loaded in main auth flow) @@ -2422,7 +2443,9 @@ class MCPRequestHandler: # not collapse to allow-all (None); key/JWT auth keeps its prior allow-all-on-error. Both # keyless_source AND the marker are needed: each source resolves through an UNMARKED auth, so # without keyless_source a fault under a source returns None and wins the union as allow-all. - deny_all = unreadable_entitlement or keyless_source or _is_mcp_admitted_user_subject(user_api_key_auth) + deny_all: Final = ( + unreadable_entitlement or keyless_source or is_mcp_admitted_user_subject(user_api_key_auth) + ) return [] if deny_all else None @staticmethod @@ -2557,7 +2580,7 @@ class MCPRequestHandler: global_mcp_server_manager, ) from litellm.proxy.auth.auth_checks import ( - _get_mcp_server_ids_from_access_groups, + get_mcp_server_ids_from_access_groups, ) from litellm.proxy.proxy_server import ( prisma_client, @@ -2565,7 +2588,7 @@ class MCPRequestHandler: user_api_key_cache, ) - raw_server_ids: Final = await _get_mcp_server_ids_from_access_groups( + raw_server_ids: Final = await get_mcp_server_ids_from_access_groups( access_group_ids=user_api_key_auth.access_group_ids or [], prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -2647,7 +2670,7 @@ class MCPRequestHandler: ) # Get MCP servers from access groups - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( key_object_permission.mcp_access_groups or [], requires_fresh_policy=user_api_key_auth.requires_fresh_policy, ) @@ -2763,7 +2786,7 @@ class MCPRequestHandler: return set(team_access_group_servers) if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []): return set(global_mcp_server_manager.get_registry().keys()) - legacy_access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + legacy_access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( object_permissions.mcp_access_groups or [], requires_fresh_policy=requires_fresh_policy, ) @@ -2797,7 +2820,7 @@ class MCPRequestHandler: """ try: from litellm.proxy.auth.auth_checks import ( - _get_mcp_server_ids_from_access_groups, + get_mcp_server_ids_from_access_groups, get_team_object, ) from litellm.proxy.proxy_server import ( @@ -2825,7 +2848,7 @@ class MCPRequestHandler: # pinned to a single team_id, but a keyless admitted identity (no team_id) unions # across all of its teams and would otherwise inherit a blocked team's MCP grants. return [] - team_access_group_servers: Final = await _get_mcp_server_ids_from_access_groups( + team_access_group_servers: Final = await get_mcp_server_ids_from_access_groups( access_group_ids=team_obj.access_group_ids or [], prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -2968,7 +2991,7 @@ class MCPRequestHandler: # Expand names/aliases to canonical server IDs (consistent with key/team/end-user path) direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []) - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( object_permissions.mcp_access_groups or [], requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) @@ -3073,7 +3096,7 @@ class MCPRequestHandler: direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permission.mcp_servers or []) # Get MCP servers from access groups - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( object_permission.mcp_access_groups or [], requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy), ) @@ -3212,7 +3235,7 @@ class MCPRequestHandler: direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []) fresh: Final = bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy) - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( object_permissions.mcp_access_groups or [], requires_fresh_policy=fresh, ) @@ -3318,7 +3341,7 @@ class MCPRequestHandler: return False object_permission: Final = user_api_key_auth.object_permission credential_scoped: Final = ( - not _is_mcp_admitted_user_subject(user_api_key_auth) + not is_mcp_admitted_user_subject(user_api_key_auth) and object_permission is not None and object_permission.mcp_servers is not None ) @@ -3566,7 +3589,7 @@ class MCPRequestHandler: expanded_direct_servers: Final = global_mcp_server_manager.expand_permission_list( obj_perm.mcp_servers or [] ) - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups( obj_perm.mcp_access_groups or [], requires_fresh_policy=user_api_key_auth.requires_fresh_policy, ) @@ -3696,7 +3719,7 @@ class MCPRequestHandler: return server_ids @staticmethod - async def _get_mcp_servers_from_access_groups( + async def get_mcp_servers_from_access_groups( access_groups: list[str], *, requires_fresh_policy: bool = False, @@ -3735,6 +3758,8 @@ class MCPRequestHandler: verbose_logger.warning("Failed to get MCP servers from access groups: %s", e) return [] + _get_mcp_servers_from_access_groups = get_mcp_servers_from_access_groups + @staticmethod async def get_mcp_access_groups( user_api_key_auth: UserAPIKeyAuth | None = None, @@ -3863,5 +3888,5 @@ class MCPRequestHandler: """ Extract and parse the x-mcp-access-groups header from an ASGI scope. """ - headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope) + headers: Final = MCPRequestHandler.safe_get_headers_from_scope(scope) return MCPRequestHandler.get_mcp_access_groups_from_headers(headers) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 77bdbd26b35..9f393f25e49 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -14,8 +14,9 @@ from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS -from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage] # reuse the encrypted credential's format discriminator +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports + _V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage] # reuse the encrypted credential's format discriminator ) from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -99,7 +100,7 @@ async def _opaque_bearer_is_gateway_credential(token: str) -> bool: user_api_key_cache, ) - if is_envelope(token) or is_refresh_envelope(token) or token.startswith(_V2_GCM_PREFIX): + if is_envelope(token) or is_refresh_envelope(token) or token.startswith(V2_GCM_PREFIX): return True try: if ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token) is not None: @@ -283,12 +284,15 @@ async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResol return _ResolvedKey(key_hash=key_hash, key=key_obj) -async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None": +async def reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None": """``None`` when the user is live, else the precise failure ``load_active_user_by_id`` found.""" loaded: Final = await load_active_user_by_id(user_id) return loaded if isinstance(loaded, str) else None +_reload_active_user_by_id: Final = reload_active_user_by_id + + UserRowSource = Literal["cache", "database"] @@ -400,12 +404,12 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol return "no_active_key" return None case "user_id": - return await _reload_active_user_by_id(identity.subject) + return await reload_active_user_by_id(identity.subject) case _: assert_never(identity.subject_type) -async def _extract_user_id_from_request(request: Request) -> str | None: +async def extract_user_id_from_request(request: Request) -> str | None: """Resolve the caller for identity binding without granting credential-write permission.""" from litellm.proxy.auth.handle_jwt import JWTIdentity # noqa: PLC0415 # proxy import cycle @@ -415,6 +419,9 @@ async def _extract_user_id_from_request(request: Request) -> str | None: return _active_key_user_id(resolved) if resolved is not None else None +_extract_user_id_from_request: Final = extract_user_id_from_request + + async def authorize_oauth_credential_request(request: Request, server_id: str) -> str | None: from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle @@ -448,13 +455,13 @@ async def can_store_oauth_credential(request: Request, auth: "UserAPIKeyAuth", s ) from litellm.proxy.auth.route_checks import RouteChecks # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle - _run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action + run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action ) write_route: Final = f"/v1/mcp/server/{server_id}/oauth-user-credential" try: RouteChecks.is_virtual_key_allowed_to_call_route(route=write_route, valid_token=auth, request=request) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=auth, request=request, request_data={}, @@ -646,7 +653,7 @@ class _BridgeMintReady: keys: "EnvelopeKeys" -def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: +def bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: """Map a bridge-mint failure value to its token-endpoint response: one place, RFC 6749 §5.2 shape (top-level ``error``, no-store headers) for every case, with a status truthful about where the failure is. The caller's request is 400, a transient gateway outage is 503, a gateway @@ -731,6 +738,9 @@ def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: ) +_bridge_mint_error_response: Final = bridge_mint_error_response + + def _key_resolution_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError: """Lift an identity-resolution failure into the mint taxonomy, preserving origin so the status stays truthful: the caller's missing credential is 400, a transient DB outage is 503, and a gateway that @@ -759,7 +769,7 @@ def _upstream_rejection_to_mint_error(rejection: _UpstreamGrantRejection) -> _Br assert_never(rejection) -async def _prepare_bridge_mint( +async def prepare_bridge_mint( request: Request, mcp_server: MCPServer, bridge_identity: "_BridgeAuthorizationCode | None" = None, @@ -820,6 +830,9 @@ async def _prepare_bridge_mint( return _BridgeMintReady(identity=identity, keys=keys) +_prepare_bridge_mint: Final = prepare_bridge_mint + + @dataclass(frozen=True, slots=True) class _BridgeRefreshReady: """A validated refresh request: the identity+keys to mint the renewed pair under, the upstream refresh @@ -854,7 +867,7 @@ def _refresh_key_failure_to_mint_error(failure: _KeyResolutionFailure) -> _Bridg assert_never(failure) -async def _prepare_bridge_refresh( +async def prepare_bridge_refresh( mcp_server: MCPServer, refresh_value: str | None ) -> "_BridgeRefreshReady | _BridgeMintError": """Phase 1 for the refresh_token grant, BEFORE the upstream exchange: open the client's refresh @@ -891,7 +904,10 @@ async def _prepare_bridge_refresh( ) -def _finish_bridge_mint( +_prepare_bridge_refresh: Final = prepare_bridge_refresh + + +def finish_bridge_mint( ready: "_BridgeMintReady", mcp_server: MCPServer, token_response: object, now: datetime ) -> "JSONResponse | _BridgeMintError": """Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held access envelope @@ -933,6 +949,9 @@ def _finish_bridge_mint( return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS) +_finish_bridge_mint: Final = finish_bridge_mint + + def _upstream_refresh_credential(token_response: object) -> "RefreshCredential | None": """Extract the upstream refresh grant from a token response, or ``None`` when there is none to seal. Each field is isinstance-checked so nothing untyped reaches the refresh envelope; ``refresh_expires_in`` diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py index aece755e4c4..2181e92baa9 100644 --- a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py @@ -83,12 +83,15 @@ def _oauth_token_error(code: str, status: int = 400) -> JSONResponse: return JSONResponse(status_code=status, content={"error": code}, headers=TOKEN_NO_CACHE_HEADERS) -def _user_id_from_session_cookie(request: Request) -> str | None: +def user_id_from_session_cookie(request: Request) -> str | None: """Return user_id from the UI ``token`` cookie, or None if missing/invalid.""" user_id, _ = _session_identity_from_cookie(request) return user_id +_user_id_from_session_cookie: Final = user_id_from_session_cookie + + def _session_identity_from_cookie(request: Request) -> tuple[str | None, str | None]: """Return ``(user_id, session_key)`` from the UI ``token`` cookie (HS256-signed with ``master_key``), or ``(None, None)`` if missing/invalid. diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index 294bb6ac80f..d2d6291233e 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -844,7 +844,7 @@ async def get_filtered_server_tools( listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id) if params is None: page = ListToolsResult( - tools=await global_mcp_server_manager._get_tools_from_server( + tools=await global_mcp_server_manager.get_tools_from_server( server=server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 3b442795c55..e2d26054c80 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -33,13 +33,14 @@ from litellm.proxy._types import ( NewMCPServerRequest, UpdateMCPServerRequest, ) -from litellm.proxy.common_utils.encrypt_decrypt_utils import ( +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports SecretMapDecodeError, - _get_salt_key, + _get_salt_key, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export decode_secret_map, decrypt_value_helper, encrypt_secret_map, encrypt_value_helper, + get_salt_key, ) from litellm.proxy.utils import PrismaClient from litellm.repositories.config_repository import ConfigRepository @@ -464,7 +465,7 @@ def _prepare_mcp_server_data( blob_value = credentials.pop(te_field, None) if blob_value is not None and te_field not in data_dict: data_dict[te_field] = blob_value - data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=_get_salt_key()) + data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=get_salt_key()) data_dict["credentials"] = safe_dumps( _bind_submitted_oauth_client(data_dict["credentials"], data.issuer, data.url) if not exclude_unset and data.auth_type == "oauth2" @@ -1474,7 +1475,7 @@ async def upsert_mcp_server_oauth_client_credentials( same way regardless of which store a server's client came from.""" from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - encrypted: Final = encrypt_credentials(credentials=MCPCredentials(**credentials), encryption_key=_get_salt_key()) + encrypted: Final = encrypt_credentials(credentials=MCPCredentials(**credentials), encryption_key=get_salt_key()) blob: Final = safe_dumps(encrypted) await _oauth_client_table_actions(prisma_client).upsert( where={"server_id": server_id}, @@ -1613,11 +1614,14 @@ def _parse_oauth_payload(decoded: str | None) -> OAuthCredentialPayload | None: return None -def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None: +def decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None: """Return the OAuth2 payload dict held in ``stored``, else ``None``.""" return _parse_oauth_payload(_decode_user_credential(stored)) +_decode_oauth_payload: Final = decode_oauth_payload + + async def rotate_mcp_user_credentials_master_key(prisma_client: PrismaClient, new_master_key: str): """Re-encrypt every ``LiteLLM_MCPUserCredentials`` row with ``new_master_key``. @@ -1865,7 +1869,7 @@ async def get_user_oauth_credential( def _server_user_credential_item( row: "prisma_db_models.LiteLLM_MCPUserCredentials", ) -> MCPServerUserCredentialListItem: - oauth_payload: Final = _decode_oauth_payload(row.credential_b64) + oauth_payload: Final = decode_oauth_payload(row.credential_b64) if oauth_payload is None: return MCPServerUserCredentialListItem( user_id=row.user_id, @@ -1986,7 +1990,7 @@ async def purge_user_oauth_credentials_for_server( invalidate_token_cache is injectable for tests; it defaults to the manager's shared invalidate_user_oauth_token_cache, the single invalidation point for per-user tokens.""" rows: Final = await _db_find_user_credential_rows(prisma_client, {"server_id": server_id}) - oauth_rows: Final = [row for row in rows if _decode_oauth_payload(row.credential_b64) is not None] + oauth_rows: Final = [row for row in rows if decode_oauth_payload(row.credential_b64) is not None] if not oauth_rows: return 0 deleted_count: Final = await _user_credential_actions(prisma_client).delete_many( @@ -2201,7 +2205,7 @@ async def resolve_user_oauth_access_token( return None try: from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, mcp_per_user_token_cache, ) @@ -2248,7 +2252,7 @@ async def resolve_user_oauth_access_token( access_token: Final[str] = cred["access_token"] if prefetched_creds is None: - ttl: Final = _compute_per_user_token_ttl(server, _remaining_token_seconds(cred.get("expires_at"))) + ttl: Final = compute_per_user_token_ttl(server, _remaining_token_seconds(cred.get("expires_at"))) await mcp_per_user_token_cache.set( user_id, server_id, access_token, ttl, identity_binding_proof=cred.get("identity_binding_proof") ) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 1f38742701e..da6ce954b1d 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -24,18 +24,24 @@ from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( TokenEndpointAuthConfigError, normalize_token_endpoint_auth_method, ) -from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _bridge_mint_error_response, +from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( # noqa: F401 # legacy module exports + _bridge_mint_error_response, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export _BridgeMintReady, _BridgeRefreshReady, - _extract_user_id_from_request, - _finish_bridge_mint, - _prepare_bridge_mint, - _prepare_bridge_refresh, - _reload_active_user_by_id, + _extract_user_id_from_request, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _finish_bridge_mint, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _prepare_bridge_mint, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _prepare_bridge_refresh, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _reload_active_user_by_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export authorize_oauth_credential_request, + bridge_mint_error_response, can_store_oauth_credential, + extract_user_id_from_request, + finish_bridge_mint, oauth_authorization_uses_gateway_credential, + prepare_bridge_mint, + prepare_bridge_refresh, + reload_active_user_by_id, ) from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation from litellm.proxy._experimental.mcp_server.faults import ( @@ -91,7 +97,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.mcp import MCPAuth, MCPCredentials from litellm.types.mcp_server.mcp_server_manager import MCPServer, MCPTokenEndpointAuthMethod @@ -431,10 +440,10 @@ def _session_cookie_user_id(request: Request) -> str | None: aggregate DCR flow's verbs receive the identity as a plain value instead of parsing cookies themselves.""" from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # circular import at module load - _user_id_from_session_cookie, + user_id_from_session_cookie, ) - return _user_id_from_session_cookie(request) + return user_id_from_session_cookie(request) def _redirect_to_litellm_login(request: Request) -> RedirectResponse: @@ -637,7 +646,7 @@ async def _store_per_user_token_server_side( client even when server-side storage fails. """ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( # noqa: PLC0415 - _compute_per_user_token_ttl, + compute_per_user_token_ttl, mcp_per_user_token_cache, ) from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 @@ -693,7 +702,7 @@ async def _store_per_user_token_server_side( await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server.server_id) # Warm the Redis cache so the first subsequent MCP call is a cache hit - ttl: Final = _compute_per_user_token_ttl(server, expires_in) + ttl: Final = compute_per_user_token_ttl(server, expires_in) await mcp_per_user_token_cache.set( user_id=user_id, server_id=server.server_id, @@ -703,7 +712,7 @@ async def _store_per_user_token_server_side( ) -def _raise_if_not_oauth2(mcp_server: MCPServer) -> None: +def raise_if_not_oauth2(mcp_server: MCPServer) -> None: """Reject a server without upstream OAuth from the gateway's authorize/token/register flow. The client-forwarded token modes (``true_passthrough`` / ``oauth_delegate``) are allowed @@ -714,10 +723,10 @@ def _raise_if_not_oauth2(mcp_server: MCPServer) -> None: Authorize path with ``persist_credentials`` enabled writes nothing to the server row). """ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # circular import with mcp_server_manager at module load - _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, + UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, ) - if mcp_server.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: + if mcp_server.auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: return raise HTTPException( status_code=400, @@ -733,6 +742,9 @@ def _raise_if_not_oauth2(mcp_server: MCPServer) -> None: ) +_raise_if_not_oauth2: Final = raise_if_not_oauth2 + + def _endpoint_not_configured_detail( mcp_server: MCPServer, endpoint_label: str, @@ -919,7 +931,7 @@ async def _resolve_oauth_authorization_user( ) -> str | RedirectResponse: """Resolve the authorization subject without replacing denied credentials with cookie grants.""" from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # proxy import cycle - _user_id_from_session_cookie, + user_id_from_session_cookie, ) use_gateway_credential: Final = enforce_binding and await oauth_authorization_uses_gateway_credential(request) @@ -928,7 +940,7 @@ async def _resolve_oauth_authorization_user( ) if use_gateway_credential and request_user_id is None: return _bridge_access_denied_redirect(redirect_uri, state, mcp_server) - user_id: Final = request_user_id or _user_id_from_session_cookie(request) + user_id: Final = request_user_id or user_id_from_session_cookie(request) if user_id is None: return _redirect_to_litellm_login(request) if not await _user_can_reach_mcp_server(user_id, mcp_server.server_id): @@ -948,7 +960,7 @@ async def authorize_with_server( scope: str | None = None, ephemeral_dcr_client: "EphemeralDcrClient | None" = None, ): - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint) if not oauth_client_registration_matches( resolved_server.dcr_issuer, resolved_server.dcr_server_url, resolved_server.issuer, resolved_server.url @@ -1084,7 +1096,7 @@ async def exchange_token_with_server( scope: str | None = None, client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, ): - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) if grant_type not in ("authorization_code", "refresh_token"): raise HTTPException(status_code=400, detail="Unsupported grant_type") @@ -1131,7 +1143,7 @@ async def exchange_token_with_server( raise HTTPException(status_code=400, detail=str(exc)) from exc request_user_id: Final = ( - await _extract_user_id_from_request(request) + await extract_user_id_from_request(request) if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None else None ) @@ -1148,9 +1160,9 @@ async def exchange_token_with_server( # identity, and unwrap the real upstream refresh token BEFORE building token_data, so the exchange # sends the upstream token and never the envelope. A failure returns without touching the upstream. if is_bridge: - prepared_refresh: Final = await _prepare_bridge_refresh(resolved_server, refresh_token) + prepared_refresh: Final = await prepare_bridge_refresh(resolved_server, refresh_token) if not isinstance(prepared_refresh, _BridgeRefreshReady): - return _bridge_mint_error_response(prepared_refresh) + return bridge_mint_error_response(prepared_refresh) bridge_mint_ready = prepared_refresh.ready bridge_upstream_refresh = prepared_refresh.upstream_refresh_token bridge_upstream_scope = prepared_refresh.upstream_scope @@ -1227,9 +1239,9 @@ async def exchange_token_with_server( # Phase 1 for a bridge authorization_code mint: resolve identity (the SSO user recovered above, or # the presented litellm key) and the envelope keys BEFORE the exchange consumes the single-use code. if is_bridge: - prepared: Final = await _prepare_bridge_mint(request, resolved_server, bridge_identity) + prepared: Final = await prepare_bridge_mint(request, resolved_server, bridge_identity) if not isinstance(prepared, _BridgeMintReady): - return _bridge_mint_error_response(prepared) + return bridge_mint_error_response(prepared) bridge_mint_ready = prepared refresh_binding: Final = resolved_server.oauth_identity_binding @@ -1269,7 +1281,7 @@ async def exchange_token_with_server( "re-runs authorization_code rather than an opaque upstream error", resolved_server.server_id, ) - return _bridge_mint_error_response("invalid_refresh") + return bridge_mint_error_response("invalid_refresh") return render_token_fault(fault) token_response = response.json() @@ -1357,10 +1369,10 @@ async def exchange_token_with_server( token_response = {**token_response, "scope": refresh_request_scope} # Phase 3: seal the upstream grant into the client-held envelope; failures map through the same # OAuth-shaped response as the phase-1 preconditions. - minted: Final = _finish_bridge_mint( + minted: Final = finish_bridge_mint( bridge_mint_ready, resolved_server, token_response, datetime.now(timezone.utc) ) - return minted if isinstance(minted, JSONResponse) else _bridge_mint_error_response(minted) + return minted if isinstance(minted, JSONResponse) else bridge_mint_error_response(minted) raw_access_token: Final = token_response.get("access_token") if isinstance(token_response, dict) else None if not isinstance(raw_access_token, str) or not raw_access_token: @@ -1937,7 +1949,7 @@ async def register_client_with_server( client_redirect_uris: list[str] | None = None, client_application_type: Literal["native", "web"] | None = None, ): - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) request_base_url: Final = get_request_base_url(request) current_redirect_uri: Final = f"{request_base_url}/callback" client_facing_redirect_uris: Final = client_redirect_uris or [current_redirect_uri] @@ -2111,7 +2123,7 @@ async def authorize( mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip) if mcp_server is None: raise HTTPException(status_code=404, detail="MCP server not found") - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) # Use server's stored client_id when caller doesn't supply one. # Raise a clear error instead of passing an empty string — an empty # client_id would silently produce a broken authorization URL. @@ -2181,7 +2193,7 @@ async def token_endpoint( code_verifier=code_verifier, refresh_token=refresh_token, master_key=master_key, - reload_user=_reload_active_user_by_id, + reload_user=reload_active_user_by_id, cache=user_api_key_cache, resource=resource, mint_proxy_credential=mint_proxy_credential, @@ -2303,7 +2315,7 @@ async def introspect_endpoint(token: str = Form(...)) -> Response: return await introspect_gateway_token( token=token, master_key=master_key, - reload_user=_reload_active_user_by_id, + reload_user=reload_active_user_by_id, cache=user_api_key_cache, ) @@ -3102,7 +3114,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None): # Get the correct base URL considering X-Forwarded-* headers request_base_url: Final = get_request_base_url(request) - request_data: Final = await _read_request_body(request=request) + request_data: Final = await read_request_body(request=request) data: Final[dict] = {**request_data} client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris")) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_context.py b/litellm/proxy/_experimental/mcp_server/mcp_context.py index 11325a9f127..8ec44f65bdf 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_context.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_context.py @@ -27,16 +27,22 @@ def get_active_mcp_request_ctx() -> "ServerRequestContext | None": # Set server-side in proxy_server.py route handlers when a request arrives via # /toolset/{name}/mcp or the toolset fallback in dynamic_mcp_route. # Never populated from client-supplied headers. -_mcp_active_toolset_id: Final[ContextVar[str | None]] = ContextVar("_mcp_active_toolset_id", default=None) +mcp_active_toolset_id: Final[ContextVar[str | None]] = ContextVar("_mcp_active_toolset_id", default=None) + +_mcp_active_toolset_id: Final = mcp_active_toolset_id # Per-request merged InitializeResult.instructions; set in MCP HTTP/SSE handlers. -_mcp_gateway_initialize_instructions: Final[ContextVar[str | None]] = ContextVar( +mcp_gateway_initialize_instructions: Final[ContextVar[str | None]] = ContextVar( "_mcp_gateway_initialize_instructions", default=None ) +_mcp_gateway_initialize_instructions: Final = mcp_gateway_initialize_instructions + # Per-request scoped server name; set in MCP HTTP/SSE handlers when the path # identifies exactly one upstream server. Never populated from client-supplied headers. -_mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gateway_server_name", default=None) +mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gateway_server_name", default=None) + +_mcp_gateway_server_name: Final = mcp_gateway_server_name # Set server-side by the /mcp/proxy route. Never populated from client-supplied headers. _mcp_proxy_mode: Final[ContextVar[bool]] = ContextVar("_mcp_proxy_mode", default=False) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_debug.py b/litellm/proxy/_experimental/mcp_server/mcp_debug.py index 4691584ecbd..bf74ea99e55 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_debug.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_debug.py @@ -390,7 +390,7 @@ class MCPDebug: server_auth_type = server.auth_type break - scope_headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope) + scope_headers: Final = MCPRequestHandler.safe_get_headers_from_scope(scope) litellm_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(scope_headers) return MCPDebug.build_debug_headers( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 65bb3a59324..a0a83339d8d 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -76,10 +76,11 @@ from litellm.integrations.custom_guardrail import ( ) from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get from litellm.llms.custom_httpx.http_handler import get_async_httpx_client -from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( +from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( # noqa: F401 # legacy module exports MCPRequestHandler, MCPServerAccess, - _is_mcp_admitted_user_subject, + _is_mcp_admitted_user_subject, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_mcp_admitted_user_subject, ) from litellm.proxy._experimental.mcp_server.catalog import _configuration_identity, _DiscoveryCache, _DiscoveryKey from litellm.proxy._experimental.mcp_server.contracts import OperationContext @@ -100,10 +101,11 @@ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( MCPPerUserTokenCache, mcp_per_user_token_cache, ) -from litellm.proxy._experimental.mcp_server.oauth_utils import ( - _redact_mcp_resource_url, +from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 # legacy module exports + _redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export canonicalize_url_identity, get_byok_www_authenticate, + redact_mcp_resource_url, ) from litellm.proxy._experimental.mcp_server.outbound_credentials import ( Error, @@ -285,12 +287,14 @@ class ListedToolsCaller: # gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes. # OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the # config-YAML and DB server loaders so the two paths cannot drift on which modes trigger discovery. -_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final[tuple[MCPAuth, ...]] = ( +UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final[tuple[MCPAuth, ...]] = ( MCPAuth.oauth2, MCPAuth.true_passthrough, MCPAuth.oauth_delegate, ) +_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final = UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + _MCP_OAUTH_DISCOVERY_ON_STARTUP_ENV: Final = "LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP" _TRUE_ENV_VALUES: Final = frozenset(("1", "true", "yes", "on")) @@ -756,7 +760,7 @@ def _flow_endpoints_missing( # A configured exchange endpoint replaces discovery entirely; only a server that must # discover its token endpoint and still has none is unresolved. return token_exchange_endpoint is None and token_url is None - if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: + if auth_type not in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: return False if oauth2_flow == "client_credentials": return token_url is None @@ -926,7 +930,7 @@ def _restrict_discovery_to_corroborated_authorization_server( def _redacted_origin_list(urls: Sequence[str]) -> str: - return ", ".join(_redact_mcp_resource_url(url) or "" for url in urls) + return ", ".join(redact_mcp_resource_url(url) or "" for url in urls) def _sanitized_error_text(exc: Exception) -> str: @@ -1078,7 +1082,7 @@ def _warn_oauth_endpoints_unresolved( "(RFC 8414)", server_ref, ", ".join(unresolved), - _redact_mcp_resource_url(server_url) or "", + redact_mcp_resource_url(server_url) or "", ) return verbose_logger.warning( @@ -1107,7 +1111,7 @@ def _write_user_env_vars_cache(user_id: str, server_id: str, values: dict[str, s _user_env_vars_cache[cache_key] = (values, time.monotonic()) -def _should_strip_caller_authorization( +def should_strip_caller_authorization( mcp_server: MCPServer, raw_headers: dict[str, str] | None, user_api_key_auth: UserAPIKeyAuth | None, @@ -1171,6 +1175,9 @@ def _should_strip_caller_authorization( ) +_should_strip_caller_authorization: Final = should_strip_caller_authorization + + LITELLM_VIRTUAL_KEY_PREFIX: Final = "sk-" @@ -1272,7 +1279,7 @@ def _openapi_forwarded_extra_headers( if not mcp_server.extra_headers or not raw_headers: return None normalized_raw: Final = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} - skip_caller_authorization: Final = _should_strip_caller_authorization( + skip_caller_authorization: Final = should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -1289,7 +1296,7 @@ def _openapi_forwarded_extra_headers( return forwarded or None -def _resolve_openapi_tool_auth( +def resolve_openapi_tool_auth( mcp_server: MCPServer, mcp_auth_header: str | None, mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None, # mutable-ok: sink shape @@ -1339,6 +1346,9 @@ def _resolve_openapi_tool_auth( return None, forwarded, None +_resolve_openapi_tool_auth: Final = resolve_openapi_tool_auth + + async def _resolve_byok_mcp_auth_header( mcp_server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, @@ -1385,7 +1395,7 @@ def _catalog_auth_header( return mcp_auth_header if catalog_auth_header is ... else catalog_auth_header -def _client_forwarded_authorization_headers( +def client_forwarded_authorization_headers( mcp_server: MCPServer, oauth2_headers: dict[str, str] | None, raw_headers: dict[str, str] | None, @@ -1399,7 +1409,7 @@ def _client_forwarded_authorization_headers( paths cannot drift, mirroring the ``_should_strip_caller_authorization`` split. """ extra_headers: Final = oauth2_headers.copy() if oauth2_headers else None - if extra_headers and _should_strip_caller_authorization( + if extra_headers and should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -1408,6 +1418,9 @@ def _client_forwarded_authorization_headers( return extra_headers +_client_forwarded_authorization_headers: Final = client_forwarded_authorization_headers + + async def _materialize_auth_headers(auth: httpx2.Auth | None) -> dict[str, str] | None: """Extract the header a resolved ``httpx2.Auth`` would set, as a plain dict, or None. @@ -1472,7 +1485,7 @@ def _redacted_registry_dump(servers: dict[str, MCPServer]) -> dict[str, dict[str } -def _caller_authorization_fans_out( +def caller_authorization_fans_out( server: MCPServer, scope_servers: list[MCPServer] | None, ) -> bool: @@ -1489,6 +1502,9 @@ def _caller_authorization_fans_out( ) +_caller_authorization_fans_out: Final = caller_authorization_fans_out + + def _extract_upstream_auth_failure( exc: BaseException, ) -> tuple[int, str | None] | None: @@ -1986,7 +2002,7 @@ class MCPServerManager: manual_issuer: Final = _blank_to_none(server.issuer) manual_authorization_url: Final = _blank_to_none(server.authorization_url) manual_token_url: Final = _blank_to_none(server.token_url) - is_discovery_auth_type: Final = server.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + is_discovery_auth_type: Final = server.auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES use_issuer_anchor: Final = server.issuer_is_anchored obo_needs_discovery: Final = self._obo_needs_endpoint_discovery( server.auth_type, @@ -2276,7 +2292,7 @@ class MCPServerManager: if raw and str(raw).strip(): self._upstream_initialize_instructions_by_server_id[server.server_id] = str(raw).strip() - async def _ensure_upstream_initialize_instructions_cached(self, server: MCPServer) -> None: + async def ensure_upstream_initialize_instructions_cached(self, server: MCPServer) -> None: """ Open one upstream session and cache InitializeResult.instructions if missing. @@ -2326,7 +2342,7 @@ class MCPServerManager: raise_on_missing=False, ) extra_headers: dict[str, str] | None = dict(resolved_static_headers) if resolved_static_headers else None - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=None, extra_headers=extra_headers, @@ -2345,6 +2361,8 @@ class MCPServerManager: e, ) + _ensure_upstream_initialize_instructions_cached = ensure_upstream_initialize_instructions_cached + def get_registry(self) -> Mapping[str, MCPServer]: """ Get the registered MCP Servers from the registry and union with the config MCP Servers @@ -2458,7 +2476,7 @@ class MCPServerManager: manual_authorization_url = _blank_to_none(server_config.get("authorization_url")) manual_token_url = _blank_to_none(server_config.get("token_url")) manual_registration_url = _blank_to_none(server_config.get("registration_url")) - is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + is_discovery_auth_type = auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES obo_needs_discovery = self._obo_needs_endpoint_discovery( auth_type, server_config.get("token_exchange_endpoint"), @@ -3110,7 +3128,7 @@ class MCPServerManager: manual_authorization_url = _blank_to_none(mcp_server.authorization_url) manual_token_url = _blank_to_none(mcp_server.token_url) manual_registration_url = _blank_to_none(mcp_server.registration_url) - is_discovery_auth_type: Final = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + is_discovery_auth_type: Final = auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES token_exchange_endpoint: Final = mcp_server.token_exchange_endpoint or ( credentials_dict.get("token_exchange_endpoint") if credentials_dict else None ) @@ -3452,7 +3470,7 @@ class MCPServerManager: # applying this rule would hide almost every admitted user's OWN submitted servers. Their # submissions are theirs by authorship, and their scope comes from the per-source union. has_explicit_object_permission: Final = ( - not _is_mcp_admitted_user_subject(user_api_key_auth) + not is_mcp_admitted_user_subject(user_api_key_auth) and key_object_permission is not None and (key_object_permission.mcp_servers is not None) ) @@ -3470,7 +3488,7 @@ class MCPServerManager: the exception fallback, and applied AFTER every union (grants, operator-open, submitted) because the scope is a ceiling over the whole reachable set; a resolver fault therefore never widens a scoped bearer to the allow-all set.""" - if user_api_key_auth is None or not _is_mcp_admitted_user_subject(user_api_key_auth): + if user_api_key_auth is None or not is_mcp_admitted_user_subject(user_api_key_auth): return None return user_api_key_auth.mcp_session_resource_server_id @@ -3505,7 +3523,7 @@ class MCPServerManager: # rides the HUMAN, not the credential: an admin's session resolves the same registry their # dashboard shows (connect-page parity), bounded like an admin key by explicit # object_permission scope, the entitlement ceiling, and the session resource scope below. - is_admitted_subject: Final = _is_mcp_admitted_user_subject(user_api_key_auth) + is_admitted_subject: Final = is_mcp_admitted_user_subject(user_api_key_auth) # The key explicitly opted out of every MCP server. Return zero before # layering on allow_all_keys or submitted servers so the opt-out is absolute. @@ -3759,7 +3777,7 @@ class MCPServerManager: blocked = 0 for sid in server_ids: s = self.get_mcp_server_by_id(sid) - if s is not None and self._is_server_accessible_from_ip(s, client_ip): + if s is not None and self.is_server_accessible_from_ip(s, client_ip): allowed.append(sid) elif s is not None: blocked += 1 @@ -3774,7 +3792,7 @@ class MCPServerManager: if server is None: verbose_logger.warning("MCP Server %s not found", server_id) return [] - return list(await self._get_tools_from_server(server)) + return list(await self.get_tools_from_server(server)) except Exception as e: verbose_logger.warning("Failed to get tools from server %s: %s", server_id, e) return [] @@ -3812,7 +3830,7 @@ class MCPServerManager: try: tools: Final = list( - await self._get_tools_from_server( + await self.get_tools_from_server( server=server, mcp_auth_header=server_auth_header, user_api_key_auth=user_api_key_auth, @@ -3861,7 +3879,7 @@ class MCPServerManager: return None @staticmethod - def _extract_subject_token( + def extract_subject_token( oauth2_headers: Mapping[str, str] | None, raw_headers: Mapping[str, str] | None, user_api_key_auth: UserAPIKeyAuth | None, @@ -3878,6 +3896,8 @@ class MCPServerManager: return None return bearer + _extract_subject_token = extract_subject_token + def _obo_subject_token( self, server: MCPServer, @@ -3892,9 +3912,9 @@ class MCPServerManager: """ if server.auth_type != MCPAuth.oauth2_token_exchange: return None - return self._extract_subject_token(None, raw_headers, user_api_key_auth) + return self.extract_subject_token(None, raw_headers, user_api_key_auth) - def _build_stdio_env( + def build_stdio_env( self, server: MCPServer, raw_headers: Mapping[str, str] | None = None, @@ -3921,6 +3941,8 @@ class MCPServerManager: return resolved_env + _build_stdio_env = build_stdio_env + def _references_per_user_env_var(self, server: MCPServer) -> bool: """True when ``server.static_headers`` reference a per-user ``${NAME}`` env var. @@ -4112,7 +4134,7 @@ class MCPServerManager: Only OBO has a discovery challenge to raise; ID-JAG's failures are plain statuses whose body already names what the user has to do, so they map through ``raise_public`` as at egress. """ - subject_token: Final = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + subject_token: Final = self.extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) match server.auth_type: case MCPAuth.oauth2_token_exchange: if not self._extract_bearer_token(oauth2_headers, None): @@ -4140,7 +4162,7 @@ class MCPServerManager: ) raise_public(err) - async def _create_mcp_client( + async def create_mcp_client( self, server: MCPServer, mcp_auth_header: str | dict[str, str] | None = None, @@ -4182,7 +4204,9 @@ class MCPServerManager: elicitation_callback=(_create_elicitation_callback() if resolved_server.allow_elicitation else None), ) - async def _get_tools_from_server( + _create_mcp_client = create_mcp_client + + async def get_tools_from_server( self, server: MCPServer, mcp_auth_header: str | dict[str, str] | None = None, @@ -4212,6 +4236,8 @@ class MCPServerManager: ) return result.tools + _get_tools_from_server = get_tools_from_server + async def get_tools_page( self, server: MCPServer, @@ -4302,18 +4328,18 @@ class MCPServerManager: for_list_tools=True, ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) # token_exchange (OBO) discovery needs the caller's token too: list it with the user's own # token (mirrors the call path), not v1's deleted client_credentials fallback. Other modes # never read the inbound bearer, so leave subject_token None to avoid forwarding it. subject_token: Final = ( - self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + self.extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) if server.auth_type == MCPAuth.oauth2_token_exchange else None ) - client = await self._create_mcp_client( + client = await self.create_mcp_client( # rebind-ok: pre-existing rebinding on a rename-only line server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -4458,10 +4484,10 @@ class MCPServerManager: return None auth: Final = caller.user_api_key_auth forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None - header_env: Final = self._build_stdio_env(server, caller.raw_headers) - stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env + header_env: Final = self.build_stdio_env(server, caller.raw_headers) + stdio_env: Final = None if header_env == self.build_stdio_env(server) else header_env caller_bearer: Final = ( - self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth) + self.extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth) if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange else None ) @@ -4608,9 +4634,9 @@ class MCPServerManager: ) or None ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=headers, @@ -4657,9 +4683,9 @@ class MCPServerManager: ) or None ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=headers, @@ -4712,9 +4738,9 @@ class MCPServerManager: ) or None ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=headers, @@ -4767,9 +4793,9 @@ class MCPServerManager: ) or None ) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=headers, @@ -4820,10 +4846,10 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -4857,10 +4883,10 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) - stdio_env: Final = self._build_stdio_env(server, raw_headers) + stdio_env: Final = self.build_stdio_env(server, raw_headers) subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -4946,7 +4972,7 @@ class MCPServerManager: "MCP OAuth endpoint discovery against %s found no authorization server metadata. Attempts: %s. " "The MCP server url may be misconfigured, or the upstream may not support OAuth discovery " "(RFC 9728 / RFC 8414)", - _redact_mcp_resource_url(server_url) or "", + redact_mcp_resource_url(server_url) or "", "; ".join(attempts) if attempts else "none recorded", ) return metadata @@ -4957,7 +4983,7 @@ class MCPServerManager: *, allow_origin_fallback: bool, ) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]: - origin: Final = _redact_mcp_resource_url(server_url) or "" + origin: Final = redact_mcp_resource_url(server_url) or "" try: client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.MCP, @@ -5006,7 +5032,7 @@ class MCPServerManager: *, allow_origin_fallback: bool, ) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]: - origin: Final = _redact_mcp_resource_url(server_url) or "" + origin: Final = redact_mcp_resource_url(server_url) or "" verbose_logger.debug( "MCP OAuth discovery for %s received status error: %s", server_url, @@ -5910,14 +5936,14 @@ class MCPServerManager: } # Create MCP request object for processing - mcp_request_obj: Final = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs) + mcp_request_obj: Final = proxy_logging_obj.create_mcp_request_object_from_kwargs(pre_hook_kwargs) # Convert to LLM format for existing guardrail compatibility. # Unified guardrails read the seeded logger off the request dict and pass it # into ``apply_guardrail``, so ``@log_guardrail_information`` bridges their # evaluations itself; the ``finally`` below covers native guardrails, which # never receive it. Same seeding the pass-through routes do. - synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs) + synthetic_llm_data: Final = proxy_logging_obj.convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs) synthetic_llm_data["litellm_logging_obj"] = litellm_logging_obj try: @@ -5930,7 +5956,9 @@ class MCPServerManager: await proxy_logging_obj.enforce_mcp_server_rate_limits(user_api_key_auth, server) if modified_data: # Convert response back to MCP format and apply modifications - modified_kwargs = proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs) + modified_kwargs: Final = proxy_logging_obj.convert_mcp_hook_response_to_kwargs( + modified_data, pre_hook_kwargs + ) if modified_kwargs.get("arguments") != arguments: hook_result["arguments"] = modified_kwargs["arguments"] if modified_kwargs.get("extra_headers"): @@ -5990,7 +6018,7 @@ class MCPServerManager: } # Seeded for the same reason as in ``pre_call_tool_check``. - synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, during_hook_kwargs) + synthetic_llm_data: Final = proxy_logging_obj.convert_mcp_to_llm_format(request_obj, during_hook_kwargs) synthetic_llm_data["litellm_logging_obj"] = litellm_logging_obj # Wrapped so the bridge runs inside the task: the caller only holds the task and @@ -6063,7 +6091,7 @@ class MCPServerManager: spec: Final = to_server_spec(mcp_server) if spec is not None: await self._cred_provider.invalidate_credentials(to_subject(user_api_key_auth, subject_token), spec) - retry_client: Final = await self._create_mcp_client( + retry_client: Final = await self.create_mcp_client( server=mcp_server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, @@ -6132,7 +6160,9 @@ class MCPServerManager: MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag, ): - subject_token = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + subject_token = self.extract_subject_token( # rebind-ok: pre-existing rebinding on a rename-only line + oauth2_headers, raw_headers, user_api_key_auth + ) elif mcp_server.auth_type == MCPAuth.oauth2: if mcp_server.has_client_credentials: # For M2M OAuth servers, Authorization must come from token fetch. @@ -6143,18 +6173,20 @@ class MCPServerManager: # token, so drop the caller-forwarded Authorization (apply-if-absent would # otherwise let it shadow the resolved token). Delegate keeps it. Centralized # via _should_strip_caller_authorization to match _prepare_mcp_server_headers. - if extra_headers and _should_strip_caller_authorization( + if extra_headers and should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, ): extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER) elif mcp_server.is_client_forwarded_token: - extra_headers = _client_forwarded_authorization_headers( - mcp_server=mcp_server, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, + extra_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line + client_forwarded_authorization_headers( + mcp_server=mcp_server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) ) if mcp_server.extra_headers and raw_headers: @@ -6162,7 +6194,7 @@ class MCPServerManager: extra_headers = {} normalized_raw_headers: Final = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} - strip_caller_authorization: Final = _should_strip_caller_authorization( + strip_caller_authorization: Final = should_strip_caller_authorization( mcp_server=mcp_server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -6216,9 +6248,9 @@ class MCPServerManager: if extra_headers is not None and len(extra_headers) == 0: extra_headers = None - stdio_env: Final = self._build_stdio_env(mcp_server, raw_headers) + stdio_env: Final = self.build_stdio_env(mcp_server, raw_headers) - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=mcp_server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, @@ -6347,7 +6379,9 @@ class MCPServerManager: ) -> MCPServer: """Resolve MCP server for call_tool (prefixed name, registry, fallback).""" prefixed_tool_name: Final = add_server_prefix_to_name(name, server_name) - mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name) + mcp_server = self.get_mcp_server_from_tool_name( # rebind-ok: pre-existing rebinding on a rename-only line + prefixed_tool_name + ) resolved_by_server_name_only = False normalized_server_name: Final = normalize_server_name(server_name) @@ -6368,7 +6402,7 @@ class MCPServerManager: resolved_by_server_name_only = True break if mcp_server is None: - fallback: Final = self._get_mcp_server_from_tool_name(name) + fallback: Final = self.get_mcp_server_from_tool_name(name) if fallback is not None and (not server_name or _candidate_matches_server_name(fallback)): mcp_server = fallback if mcp_server is None: @@ -6495,7 +6529,9 @@ class MCPServerManager: subject_token: str | None = None if isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)): - subject_token = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) + subject_token = self.extract_subject_token( # rebind-ok: pre-existing rebinding on a rename-only line + oauth2_headers, raw_headers, user_api_key_auth + ) elif isinstance(spec.config, PassthroughConfig): inbound_token, forwarded_headers = take_forwarded_authorization(forwarded_headers) per_server_token: Final = passthrough_token_from_mcp_auth_header(mcp_auth_header) @@ -6637,7 +6673,7 @@ class MCPServerManager: server_name, ) - auth_header_value, openapi_forwarded_headers, upstream_credential = _resolve_openapi_tool_auth( + auth_header_value, openapi_forwarded_headers, upstream_credential = resolve_openapi_tool_auth( mcp_server=mcp_server, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, @@ -6655,21 +6691,21 @@ class MCPServerManager: async def _call_openapi_via_handler(): from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, - _request_resolved_auth_headers, + request_auth_header, + request_extra_headers, + request_resolved_auth_headers, ) - auth_token: Final = _request_auth_header.set(auth_header_value) - extra_token: Final = _request_extra_headers.set(forwarded_headers) - resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers) + auth_token: Final = request_auth_header.set(auth_header_value) + extra_token: Final = request_extra_headers.set(forwarded_headers) + resolved_token: Final = request_resolved_auth_headers.set(resolved_auth_headers) try: async with self._limit_outbound_concurrency(mcp_server): return await self._call_openapi_tool_handler(mcp_server, name, arguments, wire_compat) finally: - _request_auth_header.reset(auth_token) - _request_extra_headers.reset(extra_token) - _request_resolved_auth_headers.reset(resolved_token) + request_auth_header.reset(auth_token) + request_extra_headers.reset(extra_token) + request_resolved_auth_headers.reset(resolved_token) tasks.append(asyncio.create_task(_call_openapi_via_handler())) else: @@ -6720,7 +6756,7 @@ class MCPServerManager: # Skip OAuth2 servers that rely on user-provided tokens continue try: - tools = await self._get_tools_from_server(server) + tools = await self.get_tools_from_server(server) except MCPUpstreamAuthError as e: # Pass-through servers expect a user-supplied bearer token; # at startup we have none, so an upstream 401 is normal. @@ -6741,7 +6777,7 @@ class MCPServerManager: self.tool_name_to_mcp_server_name_mapping[original_name] = server.name self.tool_name_to_mcp_server_name_mapping[tool.name] = server.name - def _get_mcp_server_from_tool_name(self, tool_name: str) -> MCPServer | None: + def get_mcp_server_from_tool_name(self, tool_name: str) -> MCPServer | None: """ Get the MCP Server from the tool name (handles both prefixed and non-prefixed names) @@ -6786,6 +6822,8 @@ class MCPServerManager: return None + _get_mcp_server_from_tool_name = get_mcp_server_from_tool_name + async def reload_servers_from_database(self): await self.catalog.reload() @@ -6809,7 +6847,7 @@ class MCPServerManager: # Fallback if proxy_server not available return {} - def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: + def is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: """ Check if a server is accessible from the given client IP. @@ -6830,12 +6868,14 @@ class MCPServerManager: internal_networks = IPAddressUtils.parse_internal_networks(general_settings.get("mcp_internal_ip_ranges")) return IPAddressUtils.is_internal_ip(client_ip, internal_networks) + _is_server_accessible_from_ip = is_server_accessible_from_ip + def get_mcp_server_by_id(self, server_id: str, client_ip: str | None = None) -> MCPServer | None: """Get the MCP Server from the server id.""" registry: Final = self.get_registry() for server in registry.values(): if server.server_id == server_id: - if not self._is_server_accessible_from_ip(server, client_ip): + if not self.is_server_accessible_from_ip(server, client_ip): return None return server return None @@ -6961,19 +7001,19 @@ class MCPServerManager: # Pass 1: Match by alias (highest priority) for server in registry.values(): if server.alias == server_name: - if not self._is_server_accessible_from_ip(server, client_ip): + if not self.is_server_accessible_from_ip(server, client_ip): return None return server # Pass 2: Match by server_name for server in registry.values(): if server.server_name == server_name: - if not self._is_server_accessible_from_ip(server, client_ip): + if not self.is_server_accessible_from_ip(server, client_ip): return None return server # Pass 3: Match by name (lowest priority) for server in registry.values(): if server.name == server_name: - if not self._is_server_accessible_from_ip(server, client_ip): + if not self.is_server_accessible_from_ip(server, client_ip): return None return server return None @@ -6989,7 +7029,7 @@ class MCPServerManager: registry: Final = self.get_registry() if client_ip is None: return registry - return {k: v for k, v in registry.items() if self._is_server_accessible_from_ip(v, client_ip)} + return {k: v for k, v in registry.items() if self.is_server_accessible_from_ip(v, client_ip)} def _generate_stable_server_id( self, @@ -7054,7 +7094,7 @@ class MCPServerManager: if server.spec_path: spec_status, spec_error, spec_checked_at = await self._openapi_health_probes(server.spec_path).check() - return self._build_mcp_server_table(server).model_copy( + return self.build_mcp_server_table(server).model_copy( update=MappingProxyType( { "status": spec_status, @@ -7086,7 +7126,7 @@ class MCPServerManager: raise_on_missing=False, ) extra_headers: Final = dict(resolved_static_headers) if resolved_static_headers else {} - client: Final = await self._create_mcp_client( + client: Final = await self.create_mcp_client( server=server, mcp_auth_header=None, extra_headers=extra_headers, @@ -7207,7 +7247,7 @@ class MCPServerManager: verbose_logger.warning("MCP Server %s not found in registry", server_id) continue - mcp_server_table = self._build_mcp_server_table(server) + mcp_server_table = self.build_mcp_server_table(server) list_mcp_servers.append(mcp_server_table) return list_mcp_servers @@ -7220,7 +7260,7 @@ class MCPServerManager: return None return [MCPEnvVar.model_validate(env_var) for env_var in env_vars] - def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: + def build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: return LiteLLM_MCPServerTable( server_id=server.server_id, is_config=self.is_config_declared_server(server.server_id) and server.server_id not in self.registry, @@ -7275,6 +7315,8 @@ class MCPServerManager: rpm=server.rpm, ) + _build_mcp_server_table = build_mcp_server_table + async def get_all_mcp_servers_unfiltered(self) -> list[LiteLLM_MCPServerTable]: """Return all MCP servers from registry without applying access controls.""" @@ -7284,7 +7326,7 @@ class MCPServerManager: servers: Final[list[LiteLLM_MCPServerTable]] = [] for server in registry.values(): - servers.append(self._build_mcp_server_table(server)) + servers.append(self.build_mcp_server_table(server)) return servers async def get_all_mcp_servers_with_health_unfiltered( diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py b/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py index 150900e7ff2..71281b276fa 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py @@ -42,7 +42,11 @@ from typing import Final, Literal, Protocol from pydantic import JsonValue from litellm._logging import verbose_proxy_logger -from litellm.proxy._experimental.mcp_server.db import _decode_oauth_payload, decrypt_credentials +from litellm.proxy._experimental.mcp_server.db import ( # noqa: F401 # legacy module exports + _decode_oauth_payload, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + decode_oauth_payload, + decrypt_credentials, +) from litellm.proxy.utils import PrismaClient from litellm.types.mcp import MCPCredentials @@ -156,7 +160,7 @@ async def backfill_null_oauth2_flows(prisma_client: PrismaClient) -> dict[Backfi where={"server_id": {"in": server_ids}}, ) server_ids_with_oauth_tokens: Final[set[str]] = { - token_row.server_id for token_row in token_rows if _decode_oauth_payload(token_row.credential_b64) is not None + token_row.server_id for token_row in token_rows if decode_oauth_payload(token_row.credential_b64) is not None } classified: Final = tuple( diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py index 080665fda8a..6d4fefcdb6c 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py @@ -205,7 +205,7 @@ class MCPOAuth2TokenCache(InMemoryCache): mcp_oauth2_token_cache: Final = MCPOAuth2TokenCache() -def _compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> int: +def compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> int: """Compute Redis TTL for a per-user token. Uses server.token_storage_ttl_seconds when configured, capped at the token's @@ -223,6 +223,9 @@ def _compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> return MCP_PER_USER_TOKEN_DEFAULT_TTL +_compute_per_user_token_ttl: Final = compute_per_user_token_ttl + + class MCPPerUserTokenCache: """Redis-backed cache for per-user OAuth2 access tokens. diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index ccbf69d0f54..e60c4ef1c30 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -84,7 +84,7 @@ def _origin_label(scheme: str, netloc: str) -> str: return f"{scheme}://{netloc}" if netloc else f"{scheme}://" -def _redact_mcp_resource_url(url: str | None) -> str | None: +def redact_mcp_resource_url(url: str | None) -> str | None: """Reduce an MCP server URL to its origin (scheme + host + port) for logging. Everything else is dropped: userinfo (``user:pass@``), the query string, the @@ -107,6 +107,9 @@ def _redact_mcp_resource_url(url: str | None) -> str | None: return urlunsplit((parts.scheme, netloc, "", "", "")) or None +_redact_mcp_resource_url: Final = redact_mcp_resource_url + + def _resolve_proxy_base_url_env() -> str | None: global _warned_invalid_proxy_base_url configured: Final = os.environ.get("PROXY_BASE_URL", "").strip() diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 5b23695d06d..8ef573aa489 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -27,7 +27,9 @@ from litellm.proxy._experimental.mcp_server.exceptions import ( # tag-namespaced operationIds like "actions/download-job-logs-for-workflow-run" # which include '/'. Sanitize here so the same regex passes everywhere downstream. _OPENAPI_TOOL_NAME_INVALID_CHARS: Final = re.compile(r"[^a-zA-Z0-9_-]") -_OPENAPI_TOOL_NAME_MAX_LEN: Final = 128 +OPENAPI_TOOL_NAME_MAX_LEN: Final = 128 + +_OPENAPI_TOOL_NAME_MAX_LEN: Final = OPENAPI_TOOL_NAME_MAX_LEN def sanitize_openapi_tool_name(raw_name: str) -> str: @@ -41,7 +43,7 @@ def sanitize_openapi_tool_name(raw_name: str) -> str: if not raw_name: return raw_name sanitized: Final = _OPENAPI_TOOL_NAME_INVALID_CHARS.sub("_", raw_name).lower() - return sanitized[:_OPENAPI_TOOL_NAME_MAX_LEN] + return sanitized[:OPENAPI_TOOL_NAME_MAX_LEN] from litellm._logging import verbose_logger @@ -106,23 +108,31 @@ HEADERS: Final[dict[str, str]] = {} # Per-request auth header override for BYOK servers. # Set this ContextVar before calling a local tool handler to inject the user's # stored credential into the HTTP request made by the tool function closure. -_request_auth_header: contextvars.ContextVar[str | None] = contextvars.ContextVar("_request_auth_header", default=None) +request_auth_header: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar( + "_request_auth_header", default=None +) + +_request_auth_header: Final = request_auth_header # Per-request extra headers forwarded from the client request. # Populated from MCPServer.extra_headers names matched against raw request # headers in server.py before dispatching to a local/OpenAPI tool handler. -_request_extra_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar( +request_extra_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar( "_request_extra_headers", default=None ) +_request_extra_headers: Final = request_extra_headers + # Per-request headers carrying the gateway-resolved upstream credential # (stored per-user OAuth token, minted M2M token, exchanged OBO token). # Set from MCPServerManager.resolve_openapi_upstream_auth; authoritative # over every other Authorization source in _merge_openapi_tool_request_headers. -_request_resolved_auth_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar( +request_resolved_auth_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar( "_request_resolved_auth_headers", default=None ) +_request_resolved_auth_headers: Final = request_resolved_auth_headers + _request_upstream_url: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar( "_request_upstream_url", default=None ) @@ -369,7 +379,7 @@ async def _drop_credential_across_origin(request: httpx.Request) -> None: built would never be closed. """ guard: Final = credential_redirect_hook( - _request_upstream_url.get() or "", custom_credential_slot(_request_resolved_auth_headers.get()) + _request_upstream_url.get() or "", custom_credential_slot(request_resolved_auth_headers.get()) ) if guard is not None: await guard(request) @@ -382,7 +392,7 @@ def _upstream_client() -> AsyncHTTPHandler: itself, so this arm installs the same hook the MCP client uses. Both variants come from the shared cache, so a guarded call reuses its connection pool like any other. """ - if custom_credential_slot(_request_resolved_auth_headers.get()) is None: + if custom_credential_slot(request_resolved_auth_headers.get()) is None: return get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) return get_async_httpx_client( llm_provider=httpxSpecialProvider.MCP, @@ -417,20 +427,20 @@ def _merge_openapi_tool_request_headers( Header names are compared case-insensitively so different casing cannot bypass the precedence rules. """ - request_extra: Final = _request_extra_headers.get() or {} + request_extra: Final = request_extra_headers.get() or {} static: Final = static_headers or {} static_lower_names: Final = {k.lower() for k in static} effective_headers: dict[str, str] = {k: v for k, v in request_extra.items() if k.lower() not in static_lower_names} effective_headers.update(static) - override_auth: Final = _request_auth_header.get() + override_auth: Final = request_auth_header.get() if override_auth: for existing in [k for k in effective_headers if k.lower() == "authorization"]: del effective_headers[existing] effective_headers["Authorization"] = override_auth - resolved_auth_headers: Final = _request_resolved_auth_headers.get() or {} + resolved_auth_headers: Final = request_resolved_auth_headers.get() or {} for name, value in resolved_auth_headers.items(): for existing in [k for k in effective_headers if k.lower() == name.lower()]: del effective_headers[existing] @@ -622,7 +632,7 @@ def register_tools_from_openapi(spec: Mapping[str, Any], base_url: str) -> None: while unique in used_names: n += 1 suffix = f"_{n}" - unique = tool_name[: _OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix + unique = tool_name[: OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix tool_name = unique used_names.add(tool_name) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index e226b7f3fcb..6ec5bdf6af7 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -83,23 +83,31 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( classify_list_exception, outcome_wire_value, ) -from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: F401 # legacy module exports MCPServerManager, - _caller_authorization_fans_out, - _client_forwarded_authorization_headers, - _resolve_openapi_tool_auth, - _should_strip_caller_authorization, + _caller_authorization_fans_out, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _client_forwarded_authorization_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _resolve_openapi_tool_auth, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _should_strip_caller_authorization, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + caller_authorization_fans_out, + client_forwarded_authorization_headers, global_mcp_server_manager, listed_tools_caller_for, + resolve_openapi_tool_auth, + should_strip_caller_authorization, ) -from litellm.proxy._experimental.mcp_server.oauth_utils import ( - _redact_mcp_resource_url, +from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 # legacy module exports + _redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export get_byok_www_authenticate, + redact_mcp_resource_url, ) -from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, - _request_resolved_auth_headers, +from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( # noqa: F401 # legacy module exports + _request_auth_header, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _request_extra_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _request_resolved_auth_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + request_auth_header, + request_extra_headers, + request_resolved_auth_headers, ) from litellm.proxy._experimental.mcp_server.result_conversion import ( WireCompat, @@ -523,7 +531,7 @@ async def _get_allowed_mcp_servers_from_mcp_server_names( if not server_name_matched: try: - access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_server_ids = await MCPRequestHandler.get_mcp_servers_from_access_groups( [server_or_group] ) # Only include servers that the user has access to @@ -882,7 +890,7 @@ def _prepare_mcp_server_headers( # x-mcp-{alias}-authorization. The decision is computed once so BOTH the forwarding branch and # the extra_headers copy loop below honor it — otherwise a server that lists Authorization in # extra_headers would re-copy the withheld bearer from raw_headers and replay it anyway. - withhold_forwarded_authorization: Final = is_client_forwarded_mode and _caller_authorization_fans_out( + withhold_forwarded_authorization: Final = is_client_forwarded_mode and caller_authorization_fans_out( server, scope_servers ) if server.auth_type == MCPAuth.oauth2: @@ -897,7 +905,7 @@ def _prepare_mcp_server_headers( # token, so drop the caller-forwarded Authorization (apply-if-absent would # otherwise let it shadow the resolved token). Delegate keeps it. Centralized # via _should_strip_caller_authorization to match _call_regular_mcp_tool. - if extra_headers and _should_strip_caller_authorization( + if extra_headers and should_strip_caller_authorization( mcp_server=server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -905,11 +913,13 @@ def _prepare_mcp_server_headers( extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER) elif is_client_forwarded_mode: if not withhold_forwarded_authorization: - extra_headers = _client_forwarded_authorization_headers( - mcp_server=server, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, + extra_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line + client_forwarded_authorization_headers( + mcp_server=server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) ) if server.extra_headers and raw_headers: @@ -922,7 +932,7 @@ def _prepare_mcp_server_headers( # ``MCPServerManager._call_regular_mcp_tool`` so the two # code paths cannot drift on this security-sensitive choice. # See ``_should_strip_caller_authorization`` for the rules. - strip_caller_authorization: Final = _should_strip_caller_authorization( + strip_caller_authorization: Final = should_strip_caller_authorization( mcp_server=server, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, @@ -1787,7 +1797,7 @@ def _challenge_missing_token_exchange_subject( return if all(allowed.server_id != server.server_id for allowed in allowed_mcp_servers): return - if global_mcp_server_manager._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) is not None: + if global_mcp_server_manager.extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) is not None: return from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph raise_token_exchange_challenge, @@ -1978,10 +1988,12 @@ async def _execute_mcp_tool( original_tool_name = name else: # Resolve from tool name (MCP JSON-RPC or prefixed REST tool names). - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + mcp_server = global_mcp_server_manager.get_mcp_server_from_tool_name( # rebind-ok: pre-existing rebinding on a rename-only line + name + ) if mcp_server is None and requested_server is not None: for known_prefix in iter_known_server_prefixes(requested_server): - candidate = global_mcp_server_manager._get_mcp_server_from_tool_name( + candidate = global_mcp_server_manager.get_mcp_server_from_tool_name( add_server_prefix_to_name(name, known_prefix) ) if candidate is not None: @@ -2034,7 +2046,9 @@ async def _execute_mcp_tool( # Resolve the MCP server early so BYOK checks and credential injection # apply to ALL dispatch paths (local tool registry AND managed MCP server). if mcp_server is None: - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + mcp_server = global_mcp_server_manager.get_mcp_server_from_tool_name( # rebind-ok: pre-existing rebinding on a rename-only line + name + ) client_auth_header: Final = mcp_auth_header if mcp_server: @@ -2131,7 +2145,7 @@ async def _execute_mcp_tool( verbose_logger.debug("Executing local registry tool: %s", name) # The credential rides ContextVars because the tool function has its # headers baked into the closure at registration time. - auth_header_value, openapi_forwarded_headers, upstream_credential = _resolve_openapi_tool_auth( + auth_header_value, openapi_forwarded_headers, upstream_credential = resolve_openapi_tool_auth( mcp_server=mcp_server, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, @@ -2150,15 +2164,15 @@ async def _execute_mcp_tool( forwarded_headers=openapi_forwarded_headers, ) - _auth_token: Final = _request_auth_header.set(auth_header_value) - _extra_token: Final = _request_extra_headers.set(forwarded_headers) - _resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers) + _auth_token: Final = request_auth_header.set(auth_header_value) + _extra_token: Final = request_extra_headers.set(forwarded_headers) + _resolved_token: Final = request_resolved_auth_headers.set(resolved_auth_headers) try: response = await _handle_local_mcp_tool(name, arguments, wire_compat) finally: - _request_auth_header.reset(_auth_token) - _request_extra_headers.reset(_extra_token) - _request_resolved_auth_headers.reset(_resolved_token) + request_auth_header.reset(_auth_token) + request_extra_headers.reset(_extra_token) + request_resolved_auth_headers.reset(_resolved_token) # Try managed MCP server tool (the name is bare; the prefix boundary was # already resolved above against this server's registered prefixes) @@ -2618,7 +2632,7 @@ def _get_standard_logging_mcp_tool_call( server_name: str | None, session_id: str | None = None, ) -> StandardLoggingMCPToolCall: - mcp_server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name( + mcp_server: Final = global_mcp_server_manager.get_mcp_server_from_tool_name( add_server_prefix_to_name(name, server_name) if server_name else name ) namespaced_tool_name: Final = f"{server_name}/{name}" if server_name else name @@ -2632,7 +2646,7 @@ def _get_standard_logging_mcp_tool_call( namespaced_tool_name=namespaced_tool_name, mcp_session_id=session_id, mcp_auth_mode=mcp_server.auth_type, - mcp_server_resource=_redact_mcp_resource_url(mcp_server.url), + mcp_server_resource=redact_mcp_resource_url(mcp_server.url), ) else: return StandardLoggingMCPToolCall( diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 4a1ae19b2ca..ac543e1f041 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -37,7 +37,10 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( outcome_wire_value, ) from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree -from litellm.proxy._experimental.mcp_server.oauth_utils import _redact_mcp_resource_url +from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 # legacy module exports + _redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + redact_mcp_resource_url, +) from litellm.proxy._experimental.mcp_server.result_conversion import WireCompat, complete_call_tool_result from litellm.proxy._experimental.mcp_server.ui_session_utils import ( acting_user_auth, @@ -64,7 +67,10 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload from litellm.proxy.utils import ProxyLogging -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) from litellm.types.mcp import MCPAuth from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall @@ -144,7 +150,7 @@ def _known_connection_error_message(exc: BaseException, url: str | None, timeout if isinstance(exc, TimeoutError): return ( "Failed to connect to MCP server: no valid MCP response received from " - f"{_redact_mcp_resource_url(url) or 'the server'} " + f"{redact_mcp_resource_url(url) or 'the server'} " f"within {timeout_seconds:.0f}s. Check that the LiteLLM proxy can reach this URL " "from its network (DNS, egress rules, firewalls) and that the server answers MCP requests." ) @@ -223,8 +229,9 @@ if MCP_AVAILABLE: from litellm.llms.litellm_proxy.skills.skill_search import ( DEFAULT_SKILL_SEARCH_TOP_K, ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: F401 # legacy module exports + _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, ListedToolsCaller, global_mcp_server_manager, ) @@ -242,8 +249,9 @@ if MCP_AVAILABLE: filter_tools_by_key_team_permissions, fire_mcp_tool_call_failure_logging, ) - from litellm.proxy._experimental.mcp_server.server import ( - _apply_toolset_scope, + from litellm.proxy._experimental.mcp_server.server import ( # noqa: F401 # legacy module exports + _apply_toolset_scope, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + apply_toolset_scope, reject_disallowed_mcp_client, ) from litellm.proxy._experimental.mcp_server.tool_catalog_guard import ( @@ -369,7 +377,7 @@ if MCP_AVAILABLE: virtual_mcp_server_auth_headers, virtual_raw_headers, ) = _extract_mcp_headers_from_request(request, MCPRequestHandler) - virtual_oauth2_headers: Final = MCPRequestHandler._get_oauth2_headers_from_headers(request.headers) + virtual_oauth2_headers: Final = MCPRequestHandler.get_oauth2_headers_from_headers(request.headers) if tool_name == MCP_TOOL_SEARCH_TOOL_NAME: return await handle_mcp_tool_search( query=tool_arguments.get("query", ""), @@ -656,7 +664,7 @@ if MCP_AVAILABLE: if ( _server is not None and _rest_client_ip is not None - and not global_mcp_server_manager._is_server_accessible_from_ip(_server, _rest_client_ip) + and not global_mcp_server_manager.is_server_accessible_from_ip(_server, _rest_client_ip) ): raise HTTPException( status_code=403, @@ -707,7 +715,7 @@ if MCP_AVAILABLE: record_listing: bool, ) -> list[MCPTool]: return list( - await global_mcp_server_manager._get_tools_from_server( + await global_mcp_server_manager.get_tools_from_server( server=server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, @@ -856,7 +864,7 @@ if MCP_AVAILABLE: if ( _server is not None and rest_client_ip is not None - and not global_mcp_server_manager._is_server_accessible_from_ip(_server, rest_client_ip) + and not global_mcp_server_manager.is_server_accessible_from_ip(_server, rest_client_ip) ): raise HTTPException( status_code=403, @@ -953,7 +961,7 @@ if MCP_AVAILABLE: status_code=404, detail=f"Toolset '{toolset_name}' not found", ) - return await _apply_toolset_scope(user_api_key_dict, toolset.toolset_id) + return await apply_toolset_scope(user_api_key_dict, toolset.toolset_id) @router.get("/tools/list", dependencies=[Depends(user_api_key_auth)]) @catalog_operation(global_manager) @@ -1034,8 +1042,8 @@ if MCP_AVAILABLE: # Extract auth headers from request headers: Final = request.headers raw_headers_from_request: Final = dict(headers) - mcp_auth_header: Final = MCPRequestHandler._get_mcp_auth_header_from_headers(headers) - mcp_server_auth_headers: Final = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers) + mcp_auth_header: Final = MCPRequestHandler.get_mcp_auth_header_from_headers(headers) + mcp_server_auth_headers: Final = MCPRequestHandler.get_mcp_server_auth_headers_from_headers(headers) auth_contexts: Final = await build_effective_auth_contexts(user_api_key_dict) @@ -1293,7 +1301,7 @@ if MCP_AVAILABLE: if target_server is not None: user_oauth_extra_headers = await _get_user_oauth_extra_headers(target_server, user_api_key_dict) caller_oauth2_headers: Final = ( - MCPRequestHandler._get_oauth2_headers_from_headers(request.headers) + MCPRequestHandler.get_oauth2_headers_from_headers(request.headers) if target_server is not None and target_server.auth_type in _CLIENT_FORWARDED_TOKEN_AUTH_TYPES else None ) @@ -1396,9 +1404,10 @@ if MCP_AVAILABLE: # /health/tools/list -> List tools from MCP server # For these routes users will dynamically pass the MCP connection params, they don't need to be on the MCP registry ######################################################## - from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( # noqa: F401 # legacy module exports NewMCPServerRequest, - _inherit_credentials_from_existing_server, + _inherit_credentials_from_existing_server, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + inherit_credentials_from_existing_server, ) def _extract_credentials( @@ -1461,7 +1470,7 @@ if MCP_AVAILABLE: saved_origin is not None and saved_origin == preview_origin ) request: Final = ( - _inherit_credentials_from_existing_server(new_mcp_server_request) if may_inherit else new_mcp_server_request + inherit_credentials_from_existing_server(new_mcp_server_request) if may_inherit else new_mcp_server_request ) mcp_auth_header: Final = ( request.credentials.get("auth_value") @@ -1472,8 +1481,8 @@ if MCP_AVAILABLE: # when the primary x-litellm-api-key header is absent, the Authorization value is the # caller's LiteLLM key, not an upstream token, and must never be forwarded upstream. oauth2_headers: Final = ( - MCPRequestHandler._get_oauth2_headers_from_headers(headers) - if request.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES + MCPRequestHandler.get_oauth2_headers_from_headers(headers) + if request.auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) else None ) @@ -1564,7 +1573,7 @@ if MCP_AVAILABLE: instructions=request.instructions, ) - stdio_env: Final = global_mcp_server_manager._build_stdio_env(server_model, raw_headers) + stdio_env: Final = global_mcp_server_manager.build_stdio_env(server_model, raw_headers) # For M2M OAuth servers, drop the incoming Authorization header so that # resolve_mcp_auth can auto-fetch a token via client_credentials. @@ -1617,7 +1626,7 @@ if MCP_AVAILABLE: ) with anyio.fail_after(timeout_seconds): - client: Final = await global_mcp_server_manager._create_mcp_client( + client: Final = await global_mcp_server_manager.create_mcp_client( server=server_model, mcp_auth_header=mcp_auth_header, extra_headers=merged_headers, @@ -1648,7 +1657,7 @@ if MCP_AVAILABLE: async def _preview_openapi_tools(spec_path: str) -> dict: """Generate tool previews from an OpenAPI spec without creating a server.""" from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _OPENAPI_TOOL_NAME_MAX_LEN, + OPENAPI_TOOL_NAME_MAX_LEN, build_input_schema, load_openapi_spec_async, resolve_operation_params, @@ -1681,7 +1690,7 @@ if MCP_AVAILABLE: while unique in used_names: n += 1 suffix = f"_{n}" - unique = op_id[: _OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix + unique = op_id[: OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix op_id = unique used_names.add(op_id) summary = operation.get("summary", "") @@ -1738,7 +1747,7 @@ if MCP_AVAILABLE: _test_connection_operation, mcp_auth_header=staged.mcp_auth_header, oauth2_headers=staged.oauth2_headers, - raw_headers=_safe_get_request_headers(request), + raw_headers=safe_get_request_headers(request), ) @router.post("/test/tools/list", dependencies=[Depends(user_api_key_auth)]) @@ -1797,5 +1806,5 @@ if MCP_AVAILABLE: _list_tools_operation, mcp_auth_header=staged.mcp_auth_header, oauth2_headers=staged.oauth2_headers, - raw_headers=_safe_get_request_headers(request), + raw_headers=safe_get_request_headers(request), ) diff --git a/litellm/proxy/_experimental/mcp_server/sampling_handler.py b/litellm/proxy/_experimental/mcp_server/sampling_handler.py index 361d8d5ae31..e1469f2fff7 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -772,11 +772,11 @@ async def _check_model_access(model: str, user_api_key_auth: "UserAPIKeyAuth | N import litellm from litellm.proxy._types import ModelAccessDeniedProxyException from litellm.proxy.auth.auth_checks import ( - _check_team_member_model_access, can_key_call_model, can_project_access_model, can_team_access_model, can_user_call_model, + check_team_member_model_access, get_project_object, get_team_object, get_user_object, @@ -832,7 +832,7 @@ async def _check_model_access(model: str, user_api_key_auth: "UserAPIKeyAuth | N team_model_aliases=getattr(user_api_key_auth, "team_model_aliases", None), ) if _user_id and _proxy_logging_obj: - await _check_team_member_model_access( + await check_team_member_model_access( model=model, team_object=team_obj, valid_token=user_api_key_auth, diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py index 0644238071e..2f88d8c361f 100644 --- a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py +++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py @@ -117,7 +117,7 @@ class SemanticMCPToolFilter: self.tool_router = None raise - def _extract_tool_info(self, tool) -> tuple[str, str]: + def extract_tool_info(self, tool) -> tuple[str, str]: """Extract name and description from MCP tool or OpenAI function dict.""" name: str description: str @@ -133,6 +133,8 @@ class SemanticMCPToolFilter: return name, description + _extract_tool_info = extract_tool_info + def _build_router(self, tools: list) -> None: """Build semantic router with tools (MCPTool objects or OpenAI function dicts).""" from semantic_router.routers import SemanticRouter @@ -153,7 +155,7 @@ class SemanticMCPToolFilter: self._tool_map = {} for tool in tools: - name, description = self._extract_tool_info(tool) + name, description = self.extract_tool_info(tool) self._tool_map[name] = tool routes.append( @@ -187,13 +189,13 @@ class SemanticMCPToolFilter: def _has_tools_missing_from_index(self, tools: Sequence[object]) -> bool: """Allocation-free check for any named tool not yet in the semantic index.""" - return any(name and name not in self._tool_map for name in (self._extract_tool_info(t)[0] for t in tools)) + return any(name and name not in self._tool_map for name in (self.extract_tool_info(t)[0] for t in tools)) def _tools_missing_from_index(self, tools: Sequence[object]) -> Mapping[str, object]: """Map name -> tool for every named tool not yet in the semantic index.""" return { name: tool - for name, tool in ((self._extract_tool_info(t)[0], t) for t in tools) + for name, tool in ((self.extract_tool_info(t)[0], t) for t in tools) if name and name not in self._tool_map } @@ -228,7 +230,7 @@ class SemanticMCPToolFilter: if not missing: return - descriptions: Final = {name: self._extract_tool_info(tool)[1] for name, tool in missing.items()} + descriptions: Final = {name: self.extract_tool_info(tool)[1] for name, tool in missing.items()} routes: Final = [ Route( name=name, @@ -302,7 +304,7 @@ class SemanticMCPToolFilter: verbose_logger.warning("Semantic router could not be built from the request's tools") return available_tools - available_names: Final = [name for name in (self._extract_tool_info(t)[0] for t in available_tools) if name] + available_names: Final = [name for name in (self.extract_tool_info(t)[0] for t in available_tools) if name] if not available_names: return available_tools @@ -406,7 +408,7 @@ class SemanticMCPToolFilter: # names happen to be tail-compatible with the same incoming name. available_by_name: Final[dict[str, object]] = {} for tool in available_tools: - client_name, _ = self._extract_tool_info(tool) + client_name, _ = self.extract_tool_info(tool) if client_name and client_name not in available_by_name: available_by_name[client_name] = tool diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index c046fa93d54..d1ac574dc91 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -32,9 +32,10 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( +from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( # noqa: F401 # legacy module exports MCPRequestHandler, - _is_mcp_admitted_user_subject, + _is_mcp_admitted_user_subject, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_mcp_admitted_user_subject, ) from litellm.proxy._experimental.mcp_server.client_allowlist import ( MCPClientAllowlist, @@ -47,13 +48,16 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( from litellm.proxy._experimental.mcp_server.exceptions import ( MCPUpstreamAuthError, ) -from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_active_toolset_id, - _mcp_gateway_initialize_instructions, - _mcp_gateway_server_name, +from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: F401 # legacy module exports + _mcp_active_toolset_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _mcp_gateway_initialize_instructions, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _mcp_gateway_server_name, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export _mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # server-owned request mode active_mcp_request_ctx_var, get_active_mcp_request_ctx, + mcp_active_toolset_id, + mcp_gateway_initialize_instructions, + mcp_gateway_server_name, ) from litellm.proxy._experimental.mcp_server.mcp_debug import ( MCP_AUTH_DIAGNOSTICS_SCOPE_KEY, @@ -61,9 +65,9 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import ( MCPDebug, ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( - _redact_mcp_resource_url, get_passthrough_www_authenticate, get_route_relative_request_path, + redact_mcp_resource_url, well_known_root_suffix, ) from litellm.proxy._experimental.mcp_server.ui_session_utils import ( @@ -98,6 +102,8 @@ if TYPE_CHECKING: from mcp.server.session import ServerSession as _McpServerSession +_redact_mcp_resource_url: Final = redact_mcp_resource_url + _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60 # Upper bound on concurrent stateful sessions a single caller may hold. Each # `initialize` creates a session that survives until the idle timeout, so @@ -486,6 +492,7 @@ if MCP_AVAILABLE: "mcp_get_prompt", "mcp_read_resource", "raise_denied_scoped_mcp_access", + "redact_mcp_resource_url", ) from mcp.server import Server @@ -579,10 +586,10 @@ if MCP_AVAILABLE: else base_options ) updates: Final[dict[str, str]] = {} - merged: Final = _mcp_gateway_initialize_instructions.get() + merged: Final = mcp_gateway_initialize_instructions.get() if merged is not None: updates["instructions"] = merged - scoped_server_name: Final = _mcp_gateway_server_name.get() + scoped_server_name: Final = mcp_gateway_server_name.get() if scoped_server_name is not None: updates["server_name"] = scoped_server_name return opts.model_copy(update=updates) if updates else opts @@ -1028,7 +1035,7 @@ if MCP_AVAILABLE: # cancel sibling probes or 500 the gateway initialize request. await asyncio.gather( *[ - operations.global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(s) + operations.global_mcp_server_manager.ensure_upstream_initialize_instructions_cached(s) for s in allowed if s is not None ], @@ -1041,13 +1048,13 @@ if MCP_AVAILABLE: scoped_server_name = ( scoped_server.alias or scoped_server.server_name or scoped_server.name or scoped_server.server_id ) - instructions_token: Final = _mcp_gateway_initialize_instructions.set(merged) - server_name_token: Final = _mcp_gateway_server_name.set(scoped_server_name) + instructions_token: Final = mcp_gateway_initialize_instructions.set(merged) + server_name_token: Final = mcp_gateway_server_name.set(scoped_server_name) try: yield finally: - _mcp_gateway_initialize_instructions.reset(instructions_token) - _mcp_gateway_server_name.reset(server_name_token) + mcp_gateway_initialize_instructions.reset(instructions_token) + mcp_gateway_server_name.reset(server_name_token) from litellm.proxy._experimental.mcp_server.operations import ( _MCP_CREDENTIAL_REQUEST_FIELDS, @@ -1524,7 +1531,7 @@ if MCP_AVAILABLE: scope["headers"] = [(k, v) for k, v in _headers if _normalize_header_name(k) != _mcp_session_header] return False - async def _apply_toolset_scope( + async def apply_toolset_scope( user_api_key_auth: UserAPIKeyAuth, toolset_id: str, acting_user: ActingUser = acting_user_auth, @@ -1544,7 +1551,7 @@ if MCP_AVAILABLE: of its grant sources. Admins always pass. """ from litellm.proxy._types import LiteLLM_ObjectPermissionTable - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view # A key scoped to no MCP servers opts out of every MCP path. Enforce it # here too, since toolset scoping replaces mcp_servers and would otherwise @@ -1558,13 +1565,13 @@ if MCP_AVAILABLE: ) acting: Final = await acting_user(user_api_key_auth) - is_admin: Final = _user_has_admin_view(acting) + is_admin: Final = user_api_key_has_admin_view(acting) if not is_admin and toolset_id not in await granted(acting): raise HTTPException( status_code=403, detail=f"API key does not have access to toolset '{toolset_id}'.", ) - if _is_mcp_admitted_user_subject(acting): + if is_mcp_admitted_user_subject(acting): resource_server_id: Final = acting.mcp_session_resource_server_id if resource_server_id is not None and resource_server_id not in ( await operations.global_mcp_server_manager.resolve_toolset_tool_permissions( @@ -1600,6 +1607,8 @@ if MCP_AVAILABLE: ) return acting.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) + _apply_toolset_scope: Final = apply_toolset_scope + async def _toolset_server_ids(toolset_id: str) -> set[str]: return set( await operations.global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=[toolset_id]) @@ -1674,7 +1683,7 @@ if MCP_AVAILABLE: if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth): continue - if _is_mcp_admitted_user_subject(user_api_key_auth): + if is_mcp_admitted_user_subject(user_api_key_auth): raise HTTPException( status_code=401, detail="Unauthorized", @@ -2033,10 +2042,12 @@ if MCP_AVAILABLE: # Apply toolset scope if set server-side via ContextVar (set by # /toolset/{name}/mcp and /{name}/mcp route handlers in proxy_server.py). - active_toolset_id: Final = _mcp_active_toolset_id.get() + active_toolset_id: Final = mcp_active_toolset_id.get() toolset_allowed_server_ids: set[str] | None = None if active_toolset_id and user_api_key_auth is not None: - user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) + user_api_key_auth = ( # rebind-ok: pre-existing rebinding on a rename-only line + await apply_toolset_scope(user_api_key_auth, active_toolset_id) + ) toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response @@ -2380,10 +2391,12 @@ if MCP_AVAILABLE: # Apply toolset scope if set server-side via ContextVar so the # downstream probe list matches the fully-authorized server set # (mirrors the streamable HTTP handler). - active_toolset_id: Final = _mcp_active_toolset_id.get() + active_toolset_id: Final = mcp_active_toolset_id.get() toolset_allowed_server_ids: set[str] | None = None if active_toolset_id and user_api_key_auth is not None: - user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) + user_api_key_auth = ( # rebind-ok: pre-existing rebinding on a rename-only line + await apply_toolset_scope(user_api_key_auth, active_toolset_id) + ) toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response diff --git a/litellm/proxy/_experimental/mcp_server/server_resolution.py b/litellm/proxy/_experimental/mcp_server/server_resolution.py index e828d8d8ee6..31bcd897e72 100644 --- a/litellm/proxy/_experimental/mcp_server/server_resolution.py +++ b/litellm/proxy/_experimental/mcp_server/server_resolution.py @@ -19,9 +19,13 @@ class MCPServerRegistry(Protocol): def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: ... - def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: ... + def is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: ... - def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: ... + _is_server_accessible_from_ip = is_server_accessible_from_ip + + def build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: ... + + _build_mcp_server_table = build_mcp_server_table async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth) -> list[str]: ... @@ -50,7 +54,7 @@ async def resolve_mcp_server( temporary_server: Final[MCPServer | None] = await temp_lookup(server_id) if temporary_server is not None: return ResolvedMCPServer( - table=manager._build_mcp_server_table(temporary_server), + table=manager.build_mcp_server_table(temporary_server), runtime=temporary_server, source="temp", ) @@ -64,12 +68,12 @@ async def resolve_mcp_server( registry_server: Final[MCPServer | None] = ( registry_candidate if registry_candidate is not None - and (id_client_ip is None or manager._is_server_accessible_from_ip(registry_candidate, id_client_ip)) + and (id_client_ip is None or manager.is_server_accessible_from_ip(registry_candidate, id_client_ip)) else None ) if registry_server is not None: return ResolvedMCPServer( - table=manager._build_mcp_server_table(registry_server), + table=manager.build_mcp_server_table(registry_server), runtime=registry_server, source="registry", ) @@ -78,7 +82,7 @@ async def resolve_mcp_server( named_server: Final[MCPServer | None] = manager.get_mcp_server_by_name(server_id, client_ip=name_client_ip) if named_server is not None: return ResolvedMCPServer( - table=manager._build_mcp_server_table(named_server), + table=manager.build_mcp_server_table(named_server), runtime=named_server, source="registry", ) diff --git a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py index b9e25259868..9f60f0afbb6 100644 --- a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py +++ b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py @@ -125,9 +125,9 @@ async def acting_user_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth: if not is_ui_session_credential(user_api_key_auth): return user_api_key_auth - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view - if _user_has_admin_view(user_api_key_auth): + if user_api_key_has_admin_view(user_api_key_auth): return user_api_key_auth admitted: Final = await admitted_user_context(user_api_key_auth) return admitted if admitted is not None else user_api_key_auth diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 4e8880d37c5..f671e1f9098 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -459,9 +459,9 @@ class AgentRequestHandler: """ Resolve unified access group ids to agent IDs. """ - from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups + from litellm.proxy.auth.auth_checks import get_agent_ids_from_access_groups - return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only) + return await get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only) @staticmethod async def _get_agents_from_access_groups( diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index c0f34a24a0a..7545a1117a1 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -335,9 +335,9 @@ async def _resolve_daily_activity_agent_ids( user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, ) -> tuple[str, ...] | None: - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return agent_ids permitted_agent_ids: Final = await _permitted_daily_activity_agent_ids( user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py index 78af282941d..2f0bdd20a8c 100644 --- a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py +++ b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py @@ -72,7 +72,7 @@ class _MarketplaceEntry(TypedDict, total=False): category: object -async def _get_prisma_client() -> object: +async def get_prisma_client() -> object: """Get the prisma client from proxy_server.""" from litellm.proxy.proxy_server import prisma_client @@ -84,6 +84,9 @@ async def _get_prisma_client() -> object: return prisma_client +_get_prisma_client: Final = get_prisma_client + + @router.get( "/claude-code/marketplace.json", tags=["Claude Code Marketplace"], @@ -111,7 +114,7 @@ async def get_marketplace(request: Request, key: str | None = None): ``` """ try: - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() caller: Final[UserAPIKeyAuth | None] = ( await user_api_key_auth(request=request, api_key=f"Bearer {key}") if key else None @@ -328,7 +331,7 @@ async def register_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() if not re.match(r"^[a-z0-9-]+$", request.name): raise HTTPException( @@ -408,7 +411,7 @@ async def list_plugins( List of plugins with their metadata. """ try: - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() visibility: Final[SkillVisibility] = skill_visibility(user_api_key_dict) plugins: Final[Sequence[_PluginRecord]] = await ClaudeCodePluginRepository(prisma_client).table.find_many( @@ -478,7 +481,7 @@ async def get_plugin( Plugin details including source and metadata. """ try: - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique( where={"name": plugin_name} @@ -579,7 +582,7 @@ async def update_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() _validate_plugin_source(request.source) @@ -646,7 +649,7 @@ async def enable_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique( where={"name": plugin_name} @@ -695,7 +698,7 @@ async def disable_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique( where={"name": plugin_name} @@ -744,7 +747,7 @@ async def delete_plugin( try: _require_proxy_admin(user_api_key_dict) - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique( where={"name": plugin_name} diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 2911e7801f7..f2c5d066714 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -30,7 +30,10 @@ from litellm.proxy.common_request_processing import ( resolve_litellm_call_id, ) from litellm.proxy.common_utils.error_body_call_id import error_body_call_id -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_error_payload import ( LITELLM_CALL_ID_HEADER, error_status_code, @@ -145,7 +148,7 @@ async def anthropic_response( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: result: Final = await base_llm_response_processor.base_process_llm_request( @@ -319,7 +322,7 @@ async def count_tokens( litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id")) try: - request_data: Final = await _read_request_body(request=request) + request_data: Final = await read_request_body(request=request) data: Final[dict] = {**request_data} # Extract required fields diff --git a/litellm/proxy/anthropic_endpoints/gateway_endpoints.py b/litellm/proxy/anthropic_endpoints/gateway_endpoints.py index 42ed6c425ab..4c4155379cd 100644 --- a/litellm/proxy/anthropic_endpoints/gateway_endpoints.py +++ b/litellm/proxy/anthropic_endpoints/gateway_endpoints.py @@ -52,7 +52,10 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token i from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles from litellm.proxy.anthropic_endpoints.endpoints import anthropic_response, count_tokens from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.http_parsing_utils import _safe_set_request_parsed_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_set_request_parsed_body, +) from litellm.proxy.management_endpoints.sso_helper_utils import CLI_SSO_SESSIONS_TARGET from litellm.proxy.management_endpoints.ui_sso import CliSsoTeamDetail from litellm.types.llms.base import LiteLLMBaseModel @@ -449,7 +452,7 @@ async def managed_settings(request: Request) -> Response: async def _skip_otlp_body_parsing(request: Request) -> None: - _safe_set_request_parsed_body(request=request, parsed_body={}) + safe_set_request_parsed_body(request=request, parsed_body={}) _OTLP_AUTHENTICATED: Final = (Depends(_skip_otlp_body_parsing), *_AUTHENTICATED) diff --git a/litellm/proxy/anthropic_endpoints/skills_endpoints.py b/litellm/proxy/anthropic_endpoints/skills_endpoints.py index ab513b1acc4..2c9c8efcdd8 100644 --- a/litellm/proxy/anthropic_endpoints/skills_endpoints.py +++ b/litellm/proxy/anthropic_endpoints/skills_endpoints.py @@ -170,7 +170,7 @@ async def create_skill( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -300,7 +300,7 @@ async def list_skills( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -397,7 +397,7 @@ async def get_skill( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -496,7 +496,7 @@ async def delete_skill( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index bea62f40c76..1435033e1be 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -96,9 +96,11 @@ from litellm.proxy.auth.model_access_denied import ( from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import publish_auth_cache_invalidation from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec -from litellm.proxy.common_utils.http_parsing_utils import ( - _safe_get_request_headers, - _safe_get_request_query_params, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_query_params, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, + safe_get_request_query_params, ) from litellm.proxy.common_utils.model_listing_utils import alias_map from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time @@ -463,7 +465,7 @@ def _get_router_zero_cost_cache(llm_router: Router) -> dict[str, bool] | None: return cache if isinstance(cache, dict) else None -def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None) -> bool: +def is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None) -> bool: """ Check if a model has zero cost (no configured pricing). @@ -582,6 +584,9 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None return True +_is_model_cost_zero: Final = is_model_cost_zero + + _NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({}) _TEAM_GRANT_RELATIONS: Final[Mapping[str, object]] = MappingProxyType({"litellm_model_table": True}) @@ -974,7 +979,7 @@ def route_skips_budget_checks(route: str) -> bool: def request_skips_budget_checks(route: str, model: str | list[str] | None, llm_router: Router | None) -> bool: - return route_skips_budget_checks(route=route) or _is_model_cost_zero(model=model, llm_router=llm_router) + return route_skips_budget_checks(route=route) or is_model_cost_zero(model=model, llm_router=llm_router) async def common_checks( @@ -1017,8 +1022,8 @@ async def common_checks( _model: Final[str | list[str] | None] = get_model_from_request( request_data=request_body, route=route, - request_headers=_safe_get_request_headers(request=request), - request_query_params=_safe_get_request_query_params(request=request), + request_headers=safe_get_request_headers(request=request), + request_query_params=safe_get_request_query_params(request=request), llm_router=llm_router, request=request, team_id=valid_token.team_id if valid_token is not None else None, @@ -1073,7 +1078,7 @@ async def common_checks( except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: raise - if not await _key_access_group_grants_model( + if not await key_access_group_grants_model( model=_model, valid_token=valid_token, team_object=team_object, @@ -1085,7 +1090,7 @@ async def common_checks( # 2.2. If team member has per-member model scope, enforce it if _model and team_object and valid_token and valid_token.user_id: with tracer.trace("litellm.proxy.auth.common_checks.check_team_member_model_access"): - await _check_team_member_model_access( + await check_team_member_model_access( model=_model, team_object=team_object, valid_token=valid_token, @@ -1122,7 +1127,7 @@ async def common_checks( managed_models: Final = (managed_policy.object_permission or MappingProxyType({})).get("models", ()) if not isinstance(managed_models, (list, tuple)) or not managed_models: raise HTTPException(403, "This agent has no model grants") - _can_object_call_model( + can_object_call_model( model=_resolve_team_alias( _model, team_model_aliases_for_auth_check(valid_token), valid_token.team_id, llm_router ), @@ -1261,7 +1266,7 @@ async def common_checks( budget_check_coros: Final = tuple( coro for coro in ( - _team_max_budget_check( + team_max_budget_check( team_object=team_object, proxy_logging_obj=proxy_logging_obj, valid_token=valid_token, @@ -1273,7 +1278,7 @@ async def common_checks( proxy_logging_obj=proxy_logging_obj, valid_token=valid_token, ), - _organization_max_budget_check( + organization_max_budget_check( valid_token=valid_token, team_object=team_object, prisma_client=prisma_client, @@ -1305,7 +1310,7 @@ async def common_checks( team_membership=loaded_team_membership, team_membership_loaded=team_membership_loaded, ), - _check_end_user_budget(end_user_obj=end_user_object, route=route) + check_end_user_budget(end_user_obj=end_user_object, route=route) if end_user_object is not None and end_user_object.litellm_budget_table is not None else None, ) @@ -1381,7 +1386,7 @@ def effective_user_role(user_role: str | None) -> LitellmUserRoles: return LitellmUserRoles.INTERNAL_USER -def _get_user_role( +def get_user_role( user_obj: LiteLLM_UserTable | None, ) -> LitellmUserRoles | None: if user_obj is None: @@ -1389,6 +1394,9 @@ def _get_user_role( return effective_user_role(user_obj.user_role) +_get_user_role: Final = get_user_role + + def _is_api_route_allowed( route: str, request: Request, @@ -1399,12 +1407,12 @@ def _is_api_route_allowed( """ - Route b/w api token check and normal token check """ - _user_role: Final = _get_user_role(user_obj=user_obj) + _user_role: Final = get_user_role(user_obj=user_obj) if valid_token is None: raise Exception("Invalid proxy server token passed. valid_token=None.") - if not _is_user_proxy_admin(user_obj=user_obj): # if non-admin + if not is_user_proxy_admin(user_obj=user_obj): # if non-admin RouteChecks.non_proxy_admin_allowed_routes_check( user_obj=user_obj, _user_role=_user_role, @@ -1416,7 +1424,7 @@ def _is_api_route_allowed( return True -def _is_user_proxy_admin(user_obj: LiteLLM_UserTable | None): +def is_user_proxy_admin(user_obj: LiteLLM_UserTable | None) -> bool: if user_obj is None: return False @@ -1426,6 +1434,9 @@ def _is_user_proxy_admin(user_obj: LiteLLM_UserTable | None): return False +_is_user_proxy_admin: Final = is_user_proxy_admin + + def _allowed_routes_check(user_route: str, allowed_routes: list) -> bool: """ Return if a user is allowed to access route. Helper function for `allowed_routes_check`. @@ -1723,7 +1734,7 @@ async def _apply_default_budget_to_end_user( return end_user_obj.model_copy(update=MappingProxyType({"litellm_budget_table": default_budget})) -async def _check_end_user_budget( +async def check_end_user_budget( end_user_obj: LiteLLM_EndUserTable, route: str, ) -> None: @@ -1765,6 +1776,9 @@ async def _check_end_user_budget( ) +_check_end_user_budget: Final = check_end_user_budget + + #: Columns whose non-null value makes an end-user row restrict something auth enforces. ``blocked`` #: is separate: it restricts when true rather than when merely set. _RESTRICTED_COLUMNS: Final = ("budget_id", "allowed_model_region", "default_model", "object_permission_id") @@ -2900,12 +2914,12 @@ async def _cache_management_object( @with_service_target(AUTH_OBJECTS_TARGET) -async def _cache_team_object( +async def cache_team_object( team_id: str, team_table: LiteLLM_TeamTableCachedObj, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None, -): +) -> None: ## CACHE REFRESH TIME! team_table.last_refreshed_at = time.time() @@ -2957,6 +2971,9 @@ async def _cache_team_object( await _invalidate_usage_cache_entry(usage_cache, alias_key, redis_shared=redis_shared, stale="team alias") +_cache_team_object: Final = cache_team_object + + @with_service_target(SPEND_COUNTERS_TARGET) async def _invalidate_usage_cache_entry( usage_cache: DualCache | None, @@ -3137,18 +3154,18 @@ async def delete_cache_team_object( await publish_auth_cache_invalidation(cache_key=key) -async def _cache_key_object( +async def cache_key_object( hashed_token: str, user_api_key_obj: UserAPIKeyAuth, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None, -): +) -> None: key: Final = hashed_token ## CACHE REFRESH TIME user_api_key_obj.last_refreshed_at = time.time() - cached_key_obj: Final = _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_obj) + cached_key_obj: Final = copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_obj) await _cache_management_object( key=key, value=cached_key_obj, @@ -3158,12 +3175,15 @@ async def _cache_key_object( ) +_cache_key_object: Final = cache_key_object + + @with_service_target(AUTH_OBJECTS_TARGET) -async def _delete_cache_key_object( +async def delete_cache_key_object( hashed_token: str, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None, -): +) -> None: """ Evict one key object, best-effort, matching `delete_cache_team_object` and `delete_cache_key_objects`. @@ -3196,6 +3216,9 @@ async def _delete_cache_key_object( await publish_auth_cache_invalidation(cache_key=key) +_delete_cache_key_object: Final = delete_cache_key_object + + async def delete_cache_key_objects( hashed_tokens: Sequence[str], user_api_key_cache: UserApiKeyCache, @@ -3215,7 +3238,7 @@ async def delete_cache_key_objects( """ results: Final = await asyncio.gather( *( - _delete_cache_key_object( + delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -3339,7 +3362,7 @@ async def _get_team_object_from_user_api_key_cache( ) # save the team object to cache - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=_response, user_api_key_cache=user_api_key_cache, @@ -3357,7 +3380,7 @@ async def _get_team_object_from_user_api_key_cache( @with_service_target(AUTH_OBJECTS_TARGET) -async def _get_team_object_from_cache( +async def get_team_object_from_cache( key: str, user_api_key_cache: UserApiKeyCache, parent_otel_span: Span | None, @@ -3370,6 +3393,9 @@ async def _get_team_object_from_cache( return decoded +_get_team_object_from_cache: Final = get_team_object_from_cache + + async def get_team_object( team_id: str, prisma_client: PrismaClient | None, @@ -3395,7 +3421,7 @@ async def get_team_object( key: Final = f"team_id:{team_id}" if not check_db_only: - cached_team_obj: Final = await _get_team_object_from_cache( + cached_team_obj: Final = await get_team_object_from_cache( key=key, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -3433,12 +3459,12 @@ async def get_team_object( @with_service_target(AUTH_OBJECTS_TARGET) -async def _cache_access_object( +async def cache_access_object( access_group_id: str, access_group_table: LiteLLM_AccessGroupTable, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None = None, -): +) -> None: key: Final = f"access_group_id:{access_group_id}" await user_api_key_cache.async_set_cache( key=key, @@ -3448,12 +3474,15 @@ async def _cache_access_object( ) +_cache_access_object: Final = cache_access_object + + @with_service_target(AUTH_OBJECTS_TARGET) -async def _delete_cache_access_object( +async def delete_cache_access_object( access_group_id: str, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging | None = None, -): +) -> None: key: Final = f"access_group_id:{access_group_id}" user_api_key_cache.delete_cache(key=key) @@ -3463,6 +3492,9 @@ async def _delete_cache_access_object( await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key) +_delete_cache_access_object: Final = delete_cache_access_object + + @log_db_metrics @with_service_target(AUTH_OBJECTS_TARGET) async def get_access_object( @@ -3510,7 +3542,7 @@ async def get_access_object( _response: Final = LiteLLM_AccessGroupTable.model_validate(response.dict()) # Save to cache - await _cache_access_object( + await cache_access_object( access_group_id=access_group_id, access_group_table=_response, user_api_key_cache=user_api_key_cache, @@ -3566,7 +3598,7 @@ async def get_team_object_by_alias( # Check cache first (keyed by alias) cache_key: Final = f"team_alias:{team_alias}" - cached_team_obj: Final = await _get_team_object_from_cache( + cached_team_obj: Final = await get_team_object_from_cache( key=cache_key, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -3861,7 +3893,7 @@ class ExperimentalUIJWTToken: raise Exception(f"Invalid hash key. Hash key={hashed_token}. Decrypted token={decrypted_token}. Error: {e}") -async def _fetch_key_object_from_db_with_reconnect( +async def fetch_key_object_from_db_with_reconnect( hashed_token: str, prisma_client: PrismaClient, parent_otel_span: Span | None, @@ -3889,6 +3921,9 @@ async def _fetch_key_object_from_db_with_reconnect( ) +_fetch_key_object_from_db_with_reconnect: Final = fetch_key_object_from_db_with_reconnect + + async def _fetch_key_object_from_db_unbounded( hashed_token: str, prisma_client: PrismaClient, @@ -4031,13 +4066,13 @@ async def get_key_object( None if check_db_only else await user_api_key_cache.async_get_cache(key=key, model_type=UserAPIKeyAuth) ) if user_api_key_auth is not None: - return _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth) + return copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth) if check_cache_only: raise Exception(f"Key doesn't exist in cache + check_cache_only=True. key={key}.") # else, check db - _valid_token: Final[BaseModel | None] = await _fetch_key_object_from_db_with_reconnect( + _valid_token: Final[BaseModel | None] = await fetch_key_object_from_db_with_reconnect( hashed_token=hashed_token, prisma_client=prisma_client, parent_otel_span=parent_otel_span, @@ -4079,7 +4114,7 @@ async def get_key_object( return _response # save the key object to cache - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=_response, user_api_key_cache=user_api_key_cache, @@ -4089,7 +4124,7 @@ async def get_key_object( return _response -def _copy_user_api_key_auth_for_cache( +def copy_user_api_key_auth_for_cache( user_api_key_obj: UserAPIKeyAuth, ) -> UserAPIKeyAuth: copied_key_obj: Final = user_api_key_obj.model_copy() @@ -4100,6 +4135,9 @@ def _copy_user_api_key_auth_for_cache( return copied_key_obj +_copy_user_api_key_auth_for_cache: Final = copy_user_api_key_auth_for_cache + + @log_db_metrics @with_service_target(AUTH_OBJECTS_TARGET) async def get_object_permission( @@ -4424,7 +4462,7 @@ async def _get_resources_from_access_groups( return list(set(resources)) -async def _get_models_from_access_groups( +async def get_models_from_access_groups( access_group_ids: Sequence[str], prisma_client: DatabaseClient | None = None, user_api_key_cache: UserApiKeyCache | None = None, @@ -4443,7 +4481,10 @@ async def _get_models_from_access_groups( ) -async def _get_mcp_server_ids_from_access_groups( +_get_models_from_access_groups: Final = get_models_from_access_groups + + +async def get_mcp_server_ids_from_access_groups( access_group_ids: list[str], prisma_client: PrismaClient | None = None, user_api_key_cache: UserApiKeyCache | None = None, @@ -4464,7 +4505,10 @@ async def _get_mcp_server_ids_from_access_groups( ) -async def _get_agent_ids_from_access_groups( +_get_mcp_server_ids_from_access_groups: Final = get_mcp_server_ids_from_access_groups + + +async def get_agent_ids_from_access_groups( access_group_ids: list[str], prisma_client: PrismaClient | None = None, user_api_key_cache: UserApiKeyCache | None = None, @@ -4485,6 +4529,9 @@ async def _get_agent_ids_from_access_groups( ) +_get_agent_ids_from_access_groups: Final = get_agent_ids_from_access_groups + + def _resolve_all_team_model_sentinel_for_auth_check( models: list[str], llm_router: Router | None, @@ -4499,7 +4546,7 @@ def _resolve_all_team_model_sentinel_for_auth_check( return list(dict.fromkeys(non_sentinel_models + proxy_models)) -def _check_model_access_helper( +def check_model_access_helper( model: str, llm_router: Router | None, models: list[str], @@ -4547,6 +4594,9 @@ def _check_model_access_helper( return True +_check_model_access_helper: Final = check_model_access_helper + + def _can_object_call_model( model: str | list[str], llm_router: Router | None, @@ -4621,7 +4671,7 @@ def _can_object_call_model( ## check model access for alias + underlying model - allow if either is in allowed models for m in potential_models: - if _check_model_access_helper( + if check_model_access_helper( model=m, llm_router=llm_router, models=models, @@ -4647,6 +4697,9 @@ def _can_object_call_model( ) +can_object_call_model: Final = _can_object_call_model + + def _resolve_team_alias( model: str | list[str], team_model_aliases: Mapping[str, str] | None, @@ -4707,7 +4760,7 @@ async def _check_agent_access_group_model_access( param="model", code=status.HTTP_403_FORBIDDEN, ) - _can_object_call_model( + can_object_call_model( model=dispatched, llm_router=llm_router, models=sorted(ceiling.models), @@ -4751,7 +4804,7 @@ async def _check_agent_caller_model_access( prisma_client=prisma_client, key_model_aliases=caller_key_model_aliases, ) - await _check_team_member_model_access( + await check_team_member_model_access( model=model, team_object=caller_team, valid_token=caller_auth, @@ -5112,7 +5165,7 @@ async def can_key_call_model( """ key_models: Final = _resolve_key_models_for_auth_check(valid_token=valid_token) try: - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=key_models, @@ -5125,12 +5178,12 @@ async def can_key_call_model( # Fallback: check key's access_group_ids key_access_group_ids: Final = valid_token.access_group_ids or [] if key_access_group_ids: - models_from_groups: Final = await _get_models_from_access_groups( + models_from_groups: Final = await get_models_from_access_groups( access_group_ids=key_access_group_ids, prisma_client=prisma_client, ) if models_from_groups: - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=models_from_groups, @@ -5200,7 +5253,7 @@ async def can_key_call_resolved_model( except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: raise - if not await _key_access_group_grants_model( + if not await key_access_group_grants_model( model=model, valid_token=valid_token, team_object=team_object, @@ -5210,7 +5263,7 @@ async def can_key_call_resolved_model( raise if valid_token.user_id is not None and team_object_from_lookup: - await _check_team_member_model_access( + await check_team_member_model_access( model=model, team_object=team_object, valid_token=valid_token, @@ -5265,7 +5318,7 @@ def can_org_access_model( Returns True if the team can access a specific model. """ - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=org_object.models if org_object else [], @@ -5289,7 +5342,7 @@ async def can_team_access_model( 2. If not allowed natively, falls back to access_group_ids on the team """ try: - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=team_object.models if team_object else [], @@ -5302,12 +5355,12 @@ async def can_team_access_model( # Fallback: check team's access_group_ids team_access_group_ids: Final = (team_object.access_group_ids or []) if team_object else [] if team_access_group_ids: - models_from_groups: Final = await _get_models_from_access_groups( + models_from_groups: Final = await get_models_from_access_groups( access_group_ids=team_access_group_ids, prisma_client=prisma_client, ) if models_from_groups: - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=list(dict.fromkeys([*(team_object.models if team_object else []), *models_from_groups])), @@ -5367,7 +5420,7 @@ async def get_authorized_resources_from_key_access_groups( return list(set(authorized_resources)) -async def _key_access_group_grants_model( +async def key_access_group_grants_model( model: str | list[str], valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None, @@ -5387,7 +5440,7 @@ async def _key_access_group_grants_model( if not authorized_models: return False try: - _can_object_call_model( + can_object_call_model( model=model, llm_router=llm_router, models=authorized_models, @@ -5401,6 +5454,9 @@ async def _key_access_group_grants_model( return False +_key_access_group_grants_model: Final = key_access_group_grants_model + + def can_project_access_model( model: str | list[str], project_object: LiteLLM_ProjectTable, @@ -5412,7 +5468,7 @@ def can_project_access_model( Raises ProxyException if access is denied. """ - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=project_object.models if project_object else [], @@ -5437,7 +5493,7 @@ def can_customer_access_model( ) if team_target != name and name in (end_user_object.models or ()): return - _can_object_call_model( + can_object_call_model( model=team_target, llm_router=llm_router, models=end_user_object.models, @@ -5472,7 +5528,7 @@ async def can_user_call_model( code=status.HTTP_403_FORBIDDEN, ) - return _can_object_call_model( + return can_object_call_model( model=model, llm_router=llm_router, models=user_object.models, @@ -5764,11 +5820,11 @@ def _apply_budget_exceeded_throttle(valid_token: UserAPIKeyAuth) -> bool: return True -async def _virtual_key_max_budget_check( +async def virtual_key_max_budget_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, user_obj: LiteLLM_UserTable | None = None, -): +) -> None: """ Raises: BudgetExceededError if the token is over it's max budget. @@ -5844,6 +5900,9 @@ async def _virtual_key_max_budget_check( ) +_virtual_key_max_budget_check: Final = virtual_key_max_budget_check + + async def _virtual_key_multi_budget_check( valid_token: UserAPIKeyAuth, ): @@ -5888,11 +5947,11 @@ async def _virtual_key_multi_budget_check( ) -async def _virtual_key_soft_budget_check( +async def virtual_key_soft_budget_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, user_obj: LiteLLM_UserTable | None = None, -): +) -> None: """ Triggers a budget alert if the token is over it's soft budget. @@ -5927,6 +5986,9 @@ async def _virtual_key_soft_budget_check( ) +_virtual_key_soft_budget_check: Final = virtual_key_soft_budget_check + + def _parse_email_list(raw: str | Sequence[object] | None) -> list[str]: """Parse emails from a list or comma-separated string.""" if isinstance(raw, list): @@ -5968,11 +6030,11 @@ def _merge_budget_alert_email_configs( } -async def _virtual_key_max_budget_alert_check( +async def virtual_key_max_budget_alert_check( valid_token: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, user_obj: LiteLLM_UserTable | None = None, -): +) -> None: """ Triggers a budget alert if the token has reached EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE (default 80%) of its max budget. @@ -6050,6 +6112,9 @@ async def _virtual_key_max_budget_alert_check( ) +_virtual_key_max_budget_alert_check: Final = virtual_key_max_budget_alert_check + + TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY: Final = "team_member_max_budget_alert_emails" _TEAM_MEMBER_ALERT_CONFIG_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object]) @@ -6074,7 +6139,7 @@ def _valid_alert_threshold_config(raw_config: object) -> Mapping[str, str | Sequ ) -def _team_member_max_budget_alert_check( +def team_member_max_budget_alert_check( team_id: str, team_alias: str | None, team_metadata: Mapping[str, object] | None, @@ -6108,6 +6173,9 @@ def _team_member_max_budget_alert_check( asyncio.create_task(proxy_logging_obj.budget_alerts(type="max_budget_alert", user_info=call_info)) +_team_member_max_budget_alert_check: Final = team_member_max_budget_alert_check + + async def _check_team_member_budget( team_object: LiteLLM_TeamTable | None, user_object: LiteLLM_UserTable | None, @@ -6176,7 +6244,7 @@ async def _check_team_member_budget( if not math.isfinite(team_member_budget): return - _team_member_max_budget_alert_check( + team_member_max_budget_alert_check( team_id=team_object.team_id, team_alias=team_object.team_alias, team_metadata=team_object.metadata, @@ -6198,7 +6266,7 @@ async def _check_team_member_budget( ) -async def _check_team_member_model_access( +async def check_team_member_model_access( model: str | list[str], team_object: LiteLLM_TeamTable, valid_token: UserAPIKeyAuth, @@ -6238,7 +6306,7 @@ async def _check_team_member_model_access( member_allowed_models: Final[list[str]] = loaded_membership.litellm_budget_table.allowed_models try: - _can_object_call_model( + can_object_call_model( model=model, llm_router=llm_router, models=member_allowed_models, @@ -6260,11 +6328,14 @@ async def _check_team_member_model_access( ) -async def _team_max_budget_check( +_check_team_member_model_access: Final = check_team_member_model_access + + +async def team_max_budget_check( team_object: LiteLLM_TeamTable | None, valid_token: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging, -): +) -> None: """ Check if the team is over it's max budget. @@ -6310,6 +6381,9 @@ async def _team_max_budget_check( ) +_team_max_budget_check: Final = team_max_budget_check + + async def _team_multi_budget_check( team_object: LiteLLM_TeamTable | None, ): @@ -6588,13 +6662,13 @@ async def delete_cached_project_object( ) -async def _organization_max_budget_check( +async def organization_max_budget_check( valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None, prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, -): +) -> None: """ Check if the organization is over its max budget. @@ -6687,6 +6761,9 @@ async def _organization_max_budget_check( ) +_organization_max_budget_check: Final = organization_max_budget_check + + async def _tag_max_budget_check( request_body: dict, prisma_client: PrismaClient | None, diff --git a/litellm/proxy/auth/auth_checks_organization.py b/litellm/proxy/auth/auth_checks_organization.py index 37c025f0b2c..6820ded0eba 100644 --- a/litellm/proxy/auth/auth_checks_organization.py +++ b/litellm/proxy/auth/auth_checks_organization.py @@ -133,7 +133,7 @@ def get_user_organization_info( return _user_organizations, _user_organization_role_mapping -def _user_is_org_admin( +def user_is_org_admin( request_data: dict, user_object: LiteLLM_UserTable | None = None, ) -> bool: @@ -173,6 +173,9 @@ def _user_is_org_admin( return all(org_id in admin_org_ids for org_id in candidate_org_ids) +_user_is_org_admin: Final = user_is_org_admin + + TEAM_ORG_CONTEXT_ROUTES: Final = frozenset({"/team/update"}) # The RESTful update route carries the team id in the path. Match on the route # template so the sibling /team/ routes (which share the single-segment diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index a59a12d6807..0266131fc33 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -20,8 +20,9 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_utils import ( - _get_request_ip_address, +from litellm.proxy.auth.auth_utils import ( # noqa: F401 # legacy module exports + _get_request_ip_address, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_request_ip_address, is_invalid_virtual_key_error, mark_invalid_virtual_key_error, normalize_request_route, @@ -125,7 +126,7 @@ def _identity_log_suffix(resolved_identity: UserAPIKeyAuth | None) -> str: class UserAPIKeyAuthExceptionHandler: @staticmethod - async def _handle_authentication_error( + async def handle_authentication_error( e: Exception, request: Request, request_data: dict[str, object], @@ -180,7 +181,7 @@ class UserAPIKeyAuthExceptionHandler: ) else: # raise the exception to the caller - requester_ip: Final = _get_request_ip_address( + requester_ip: Final = get_request_ip_address( request=request, use_x_forwarded_for=general_settings.get("use_x_forwarded_for") is True, ) @@ -258,3 +259,5 @@ class UserAPIKeyAuthExceptionHandler: extra=log_extra, ) raise final_exception + + _handle_authentication_error = handle_authentication_error diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index e49fa6e9acd..edd6028f07c 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -77,7 +77,7 @@ def mark_invalid_virtual_key_error(exception: ProxyException, is_invalid_virtual return marked_exception -def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None: +def get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None: client_ip = None if use_x_forwarded_for is True and "x-forwarded-for" in request.headers: client_ip = request.headers["x-forwarded-for"] @@ -89,6 +89,9 @@ def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = return client_ip +_get_request_ip_address: Final = get_request_ip_address + + def _check_valid_ip( allowed_ips: list[str] | None, request: Request, @@ -101,7 +104,7 @@ def _check_valid_ip( return True, None # if general_settings.get("use_x_forwarded_for") is True then use x-forwarded-for - client_ip: Final = _get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) + client_ip: Final = get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) # Check if IP address is allowed if client_ip not in allowed_ips: @@ -1884,14 +1887,14 @@ def _extract_models_from_managed_resource_id( try: from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, decode_model_from_file_id, get_model_id_from_unified_batch_id, get_models_from_unified_file_id, + is_base64_encoded_unified_file_id, ) _append_model_candidates(candidates=candidates, value=decode_model_from_file_id(resource_id)) - unified_file_id: Final = _is_base64_encoded_unified_file_id(resource_id) + unified_file_id: Final = is_base64_encoded_unified_file_id(resource_id) if unified_file_id: _append_model_candidates( candidates=candidates, @@ -2180,7 +2183,7 @@ def _router_model_from_azure_route(route: str, llm_router: Router | None) -> str def _model_from_bedrock_route(route: str) -> str | None: from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( - _extract_model_from_bedrock_endpoint, + extract_model_from_bedrock_endpoint, is_bedrock_count_tokens_endpoint, ) @@ -2188,7 +2191,7 @@ def _model_from_bedrock_route(route: str) -> str | None: if is_bedrock_count_tokens_endpoint(bedrock_endpoint): return None try: - return _extract_model_from_bedrock_endpoint(bedrock_endpoint) + return extract_model_from_bedrock_endpoint(bedrock_endpoint) except ValueError: return None diff --git a/litellm/proxy/auth/authorization.py b/litellm/proxy/auth/authorization.py index cd0a7acff0a..bb498c33763 100644 --- a/litellm/proxy/auth/authorization.py +++ b/litellm/proxy/auth/authorization.py @@ -38,10 +38,10 @@ async def resolve_owned_read_scope( def can_read_team_logs(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable) -> bool: from litellm.proxy.management.teams.authz import is_team_admin from litellm.proxy.management_endpoints.common_utils import ( - _team_member_has_permission, # pyright: ignore[reportPrivateUsage] # reuse existing team permission policy + team_member_has_permission, # pyright: ignore[reportPrivateUsage] # reuse existing team permission policy ) - return is_team_admin(user_api_key_dict=auth, team_obj=team) or _team_member_has_permission( + return is_team_admin(user_api_key_dict=auth, team_obj=team) or team_member_has_permission( user_api_key_dict=auth, team_obj=team, permission=KeyManagementRoutes.SPEND_LOGS.value, diff --git a/litellm/proxy/auth/fallback_budget.py b/litellm/proxy/auth/fallback_budget.py index 9029f9d6996..5aa28a08bb8 100644 --- a/litellm/proxy/auth/fallback_budget.py +++ b/litellm/proxy/auth/fallback_budget.py @@ -39,8 +39,9 @@ from pydantic import ValidationError from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import ( - _is_model_cost_zero, # pyright: ignore[reportPrivateUsage] # the zero-cost predicate the auth-time budget checks use; no public equivalent +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _is_model_cost_zero, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_model_cost_zero, # pyright: ignore[reportPrivateUsage] # the zero-cost predicate the auth-time budget checks use; no public equivalent ) from litellm.router import Router from litellm.types.llms.base import LiteLLMBaseModel @@ -109,7 +110,7 @@ async def is_token_within_budget_for_model(*, model: str, valid_token: UserAPIKe A zero-cost fallback target is always allowed: refusing it would deny a request on spend some other model accrued, which is the same reasoning behind the auth-time bypass. """ - if _is_model_cost_zero(model=model, llm_router=llm_router): + if is_model_cost_zero(model=model, llm_router=llm_router): return True key_budget: Final = valid_token.max_budget diff --git a/litellm/proxy/auth/ip_address_utils.py b/litellm/proxy/auth/ip_address_utils.py index c2614b85016..5d32cc7eb4c 100644 --- a/litellm/proxy/auth/ip_address_utils.py +++ b/litellm/proxy/auth/ip_address_utils.py @@ -16,7 +16,10 @@ from fastapi import Request from pydantic import TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger -from litellm.proxy.auth.auth_utils import _get_request_ip_address +from litellm.proxy.auth.auth_utils import ( # noqa: F401 # legacy module exports + _get_request_ip_address, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_request_ip_address, +) # One-shot warning so operators upgrading from the prior "always trust X-Forwarded-*" # behaviour see an actionable message in their logs the first time it triggers. @@ -368,4 +371,4 @@ class IPAddressUtils: return client_ip case _HopCountUnset(): pass - return _get_request_ip_address(request, use_x_forwarded_for=use_xff) + return get_request_ip_address(request, use_x_forwarded_for=use_xff) diff --git a/litellm/proxy/auth/resolvers/store.py b/litellm/proxy/auth/resolvers/store.py index 13e24387551..0c8d05e0b94 100644 --- a/litellm/proxy/auth/resolvers/store.py +++ b/litellm/proxy/auth/resolvers/store.py @@ -8,10 +8,13 @@ from pydantic import BaseModel from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import ( - _cache_key_object, - _copy_user_api_key_auth_for_cache, - _fetch_key_object_from_db_with_reconnect, +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _copy_user_api_key_auth_for_cache, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _fetch_key_object_from_db_with_reconnect, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + cache_key_object, + copy_user_api_key_auth_for_cache, + fetch_key_object_from_db_with_reconnect, get_object_permission, ) from litellm.proxy.auth.auth_method import AuthMethod @@ -81,7 +84,7 @@ class IdentityStore: network: NetworkContext | None = None, ) -> Principal: key: Final = await self._resolve_key(hashed_token) - return self._principal_from_key( + return self.principal_from_key( key, auth_method=auth_method, network=network, @@ -108,12 +111,12 @@ class IdentityStore: cached: Final = await self._cache.async_get_cache(key=hashed_token, model_type=UserAPIKeyAuth) if cached is not None: - return _copy_user_api_key_auth_for_cache(user_api_key_obj=cached) + return copy_user_api_key_auth_for_cache(user_api_key_obj=cached) if self._check_cache_only: raise KeyNotInCacheError(hashed_token) - from_db: Final[BaseModel | None] = await _fetch_key_object_from_db_with_reconnect( + from_db: Final[BaseModel | None] = await fetch_key_object_from_db_with_reconnect( hashed_token=hashed_token, prisma_client=self._prisma, parent_otel_span=self._parent_otel_span, @@ -140,7 +143,7 @@ class IdentityStore: e, ) - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=key, user_api_key_cache=self._cache, @@ -149,7 +152,7 @@ class IdentityStore: return key @staticmethod - def _principal_from_key( + def principal_from_key( key: UserAPIKeyAuth, *, auth_method: AuthMethod, @@ -192,3 +195,5 @@ class IdentityStore: network=network or NetworkContext(), source_key=key, ) + + _principal_from_key = principal_from_key diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 4a4913fd65b..2c1643c7c42 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -15,7 +15,10 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) -from .auth_checks_organization import _user_is_org_admin +from .auth_checks_organization import ( # noqa: F401 # legacy module exports + _user_is_org_admin, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_is_org_admin, +) # Management write routes denied to PROXY_ADMIN_VIEW_ONLY. Adding a new write # endpoint to a management router REQUIRES adding it here too — the surrounding @@ -128,7 +131,7 @@ class RouteChecks: allowed_route in _AUTH_ENFORCED_PASS_THROUGH_ROUTE_GROUPS and RouteChecks.is_auth_enforced_pass_through_route( route=route, - method=RouteChecks._get_request_method(request=request), + method=RouteChecks.get_request_method(request=request), ) ): if RouteChecks.check_passthrough_route_access(route=route, user_api_key_dict=valid_token): @@ -141,7 +144,7 @@ class RouteChecks: # For llm_api_routes, also check registered pass-through endpoints ################################################ if allowed_route == "llm_api_routes": - if route == "/auto_router/session" and RouteChecks._get_request_method(request) == "GET": + if route == "/auto_router/session" and RouteChecks.get_request_method(request) == "GET": return True from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( @@ -151,7 +154,7 @@ class RouteChecks: if InitPassThroughEndpointHelpers.is_registered_pass_through_route(route=route): if RouteChecks.is_auth_enforced_pass_through_route( route=route, - method=RouteChecks._get_request_method(request=request), + method=RouteChecks.get_request_method(request=request), ): if RouteChecks.check_passthrough_route_access( route=route, user_api_key_dict=valid_token @@ -277,7 +280,7 @@ class RouteChecks: if RouteChecks.is_auth_enforced_pass_through_route( route=route, - method=RouteChecks._get_request_method(request=request), + method=RouteChecks.get_request_method(request=request), ): RouteChecks._require_auth_pass_through_access( route=route, @@ -329,7 +332,7 @@ class RouteChecks: elif ( _user_role == LitellmUserRoles.INTERNAL_USER.value and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.internal_user_routes.value) - or _user_is_org_admin(request_data=request_data, user_object=user_obj) + or user_is_org_admin(request_data=request_data, user_object=user_obj) and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.org_admin_allowed_routes.value) or _user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value and RouteChecks.check_route_access( @@ -423,7 +426,7 @@ class RouteChecks: if RouteChecks._route_matches_pattern(route=route, pattern=openai_route): return True # Check for wildcard patterns like "/containers/*" - if RouteChecks._is_wildcard_pattern(pattern=openai_route): + if RouteChecks.is_wildcard_pattern(pattern=openai_route): if RouteChecks.route_matches_wildcard_pattern(route=route, pattern=openai_route): return True @@ -549,12 +552,14 @@ class RouteChecks: return False @staticmethod - def _is_wildcard_pattern(pattern: str) -> bool: + def is_wildcard_pattern(pattern: str) -> bool: """ Check if pattern is a wildcard pattern """ return pattern.endswith("*") + _is_wildcard_pattern = is_wildcard_pattern + @staticmethod def route_matches_wildcard_pattern(route: str, pattern: str) -> bool: """ @@ -635,7 +640,7 @@ class RouteChecks: if any( RouteChecks.route_matches_wildcard_pattern(route=route, pattern=allowed_route) for allowed_route in allowed_routes - if RouteChecks._is_wildcard_pattern(pattern=allowed_route) + if RouteChecks.is_wildcard_pattern(pattern=allowed_route) ): return True @@ -653,7 +658,7 @@ class RouteChecks: return False @staticmethod - def _get_request_method(request: Request | None) -> str | None: + def get_request_method(request: Request | None) -> str | None: if request is None: return None @@ -666,6 +671,8 @@ class RouteChecks: return method.upper() + _get_request_method = get_request_method + @staticmethod def is_auth_enforced_pass_through_route(route: str, method: str | None = None) -> bool: """ @@ -829,7 +836,7 @@ class RouteChecks: return False @staticmethod - def _is_assistants_api_request(request: Request) -> bool: + def is_assistants_api_request(request: Request) -> bool: """ Returns True if `thread` or `assistant` is in the request path @@ -847,6 +854,8 @@ class RouteChecks: return True return False + _is_assistants_api_request = is_assistants_api_request + @staticmethod def is_generate_content_route(route: str) -> bool: """ diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 65f9b09a4d3..afe88fe09a3 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -41,22 +41,26 @@ from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.proxy._types import * from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_from_headers -from litellm.proxy.auth.auth_checks import ( +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports ExperimentalUIJWTToken, TeamNotFoundError, - _cache_key_object, - _can_object_call_model, - _check_end_user_budget, - _delete_cache_key_object, - _get_user_role, - _is_model_cost_zero, - _is_user_proxy_admin, - _team_member_max_budget_alert_check, - _virtual_key_max_budget_alert_check, - _virtual_key_max_budget_check, - _virtual_key_soft_budget_check, + _cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _can_object_call_model, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _check_end_user_budget, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _delete_cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_user_role, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _is_model_cost_zero, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _is_user_proxy_admin, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _team_member_max_budget_alert_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _virtual_key_max_budget_alert_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _virtual_key_max_budget_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _virtual_key_soft_budget_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + cache_key_object, can_key_call_model, + can_object_call_model, + check_end_user_budget, common_checks, + delete_cache_key_object, get_end_user_object, get_jwt_key_mapping_object, get_key_end_user_budget_id, @@ -66,11 +70,18 @@ from litellm.proxy.auth.auth_checks import ( get_team_membership, get_team_object, get_user_object, + get_user_role, + is_model_cost_zero, + is_user_proxy_admin, is_valid_fallback_model, jwt_key_mapping_cache_key, key_model_aliases_for_auth_check, resolve_and_validate_end_user_id, resolve_default_end_user_budget, + team_member_max_budget_alert_check, + virtual_key_max_budget_alert_check, + virtual_key_max_budget_check, + virtual_key_soft_budget_check, ) from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler from litellm.proxy.auth.auth_method import AuthMethod @@ -113,18 +124,25 @@ from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.team_grants import team_grants from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, - _safe_get_request_query_params, - _safe_set_request_parsed_body, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_query_params, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export is_opaque_audio_pass_through_request, populate_request_with_path_params, read_raw_json_body, + read_request_body, rewrite_request_model, + safe_get_request_headers, + safe_get_request_query_params, + safe_set_request_parsed_body, ) from litellm.proxy.common_utils.model_listing_utils import claude_code_requested_group -from litellm.proxy.common_utils.realtime_utils import _realtime_request_body +from litellm.proxy.common_utils.realtime_utils import ( # noqa: F401 # legacy module exports + _realtime_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + realtime_request_body, +) from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, end_user_cache_key, @@ -222,8 +240,8 @@ def _get_model_from_request_context( return get_model_from_request( request_data=request_data, route=route, - request_headers=_safe_get_request_headers(request=request), - request_query_params=_safe_get_request_query_params(request=request), + request_headers=safe_get_request_headers(request=request), + request_query_params=safe_get_request_query_params(request=request), llm_router=llm_router, request=request, team_id=team_id, @@ -477,7 +495,7 @@ async def _check_key_model_budget_with_fallback( llm_router=llm_router, ) if valid_token.team_models: - _can_object_call_model( + can_object_call_model( model=fallback_model, llm_router=llm_router, models=valid_token.team_models, @@ -489,9 +507,9 @@ async def _check_key_model_budget_with_fallback( except ProxyException: raise e request_data["model"] = fallback_model - _safe_set_request_parsed_body(request=request, parsed_body=request_data) - request._json = request_data - request._body = orjson.dumps(request_data) + safe_set_request_parsed_body(request=request, parsed_body=request_data) + request._json = request_data # pyright: ignore[reportPrivateUsage] # Starlette JSON cache + request._body = orjson.dumps(request_data) # pyright: ignore[reportPrivateUsage] # Starlette body cache path_params: Final = request.scope.get("path_params") if isinstance(path_params, dict) and "model" in path_params: path_params["model"] = fallback_model @@ -591,9 +609,9 @@ def _should_route_jwt_to_oauth2_override(token: str, jwt_handler: JWTHandler) -> return False -def _get_bearer_token( +def get_bearer_token( api_key: str, -): +) -> str: if api_key.startswith("Bearer "): # ensure Bearer token passed in api_key = api_key.replace("Bearer ", "") # extract the token elif api_key.startswith("Basic "): @@ -619,6 +637,9 @@ def _get_bearer_token( return api_key +_get_bearer_token: Final = get_bearer_token + + def _apply_budget_limits_to_end_user_params( end_user_params: dict, budget_info: LiteLLM_BudgetTable, @@ -675,10 +696,10 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str synthetic_scope[key] = ws_scope[key] request: Final = Request(scope=synthetic_scope) - request._url = websocket.url + request._url = websocket.url # pyright: ignore[reportPrivateUsage] # Starlette WebSocket URL storage async def return_body(): - return _realtime_request_body(model) + return realtime_request_body(model) request.body = return_body @@ -739,7 +760,7 @@ def update_valid_token_with_end_user_params(valid_token: UserAPIKeyAuth, end_use _global_spend_coordinator: Final = EventDrivenCacheCoordinator(log_prefix="[GLOBAL SPEND]") -async def _fetch_global_spend_with_event_coordination( +async def fetch_global_spend_with_event_coordination( cache_key: str, user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, @@ -766,6 +787,9 @@ async def _fetch_global_spend_with_event_coordination( ) +_fetch_global_spend_with_event_coordination: Final = fetch_global_spend_with_event_coordination + + async def get_global_proxy_spend( litellm_proxy_admin_name: str, user_api_key_cache: UserApiKeyCache, @@ -777,10 +801,12 @@ async def get_global_proxy_spend( if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget # Use event-driven coordination to prevent cache stampede cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY - global_proxy_spend = await _fetch_global_spend_with_event_coordination( - cache_key=cache_key, - user_api_key_cache=user_api_key_cache, - prisma_client=prisma_client, + global_proxy_spend = ( # rebind-ok: pre-existing rebinding on a rename-only line + await fetch_global_spend_with_event_coordination( + cache_key=cache_key, + user_api_key_cache=user_api_key_cache, + prisma_client=prisma_client, + ) ) if global_proxy_spend is not None: user_info: Final = CallInfo( @@ -824,7 +850,7 @@ def get_api_key( """ from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.common_utils.http_parsing_utils import ( - _safe_get_request_query_params, + safe_get_request_query_params, ) api_key = api_key @@ -834,7 +860,7 @@ def get_api_key( api_key = _get_bearer_token_or_received_api_key(custom_litellm_key_header) elif isinstance(api_key, str) and len(api_key) > 0: passed_in_key = api_key - api_key = _get_bearer_token(api_key=api_key) + api_key = get_bearer_token(api_key=api_key) elif isinstance(azure_api_key_header, str): passed_in_key = azure_api_key_header api_key = azure_api_key_header @@ -850,9 +876,9 @@ def get_api_key( elif ( RouteChecks.is_generate_content_route(route=route) and request is not None - and _safe_get_request_query_params(request).get("key") + and safe_get_request_query_params(request).get("key") ): - google_auth_key: Final[str] = _safe_get_request_query_params(request).get("key") or "" + google_auth_key: Final[str] = safe_get_request_query_params(request).get("key") or "" passed_in_key = google_auth_key api_key = google_auth_key elif pass_through_endpoints is not None: @@ -1388,7 +1414,7 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None: return parent_otel_span: Final = open_telemetry_logger.create_litellm_proxy_request_started_span( start_time=start_time, - headers=_safe_get_request_headers(request), + headers=safe_get_request_headers(request), ) # Under V2 the FastAPI instrumentor stamps http.route / url.path on the server # span; only the legacy logger needs these set explicitly. @@ -1413,12 +1439,12 @@ async def _read_request_body_deferring_parse_failure( """ if is_opaque_audio_pass_through_request( route=get_request_route(request=request), - content_type=_safe_get_request_headers(request=request).get("content-type", ""), + content_type=safe_get_request_headers(request=request).get("content-type", ""), ): - _safe_set_request_parsed_body(request=request, parsed_body={}) + safe_set_request_parsed_body(request=request, parsed_body={}) return {}, None try: - parsed_body: Final = await _read_request_body(request=request) + parsed_body: Final = await read_request_body(request=request) except ProxyException as parse_exception: return {}, parse_exception return populate_request_with_path_params(request_data=parsed_body, request=request), None @@ -1481,7 +1507,7 @@ async def _refresh_session_token_grants( { **valid_token.model_dump(exclude_none=True), **team_grants(team_object, team_membership, user_object.user_id), - "user_role": _get_user_role(user_object), + "user_role": get_user_role(user_object), "models": () if team_object is not None else user_models(user_object), } ) @@ -1515,7 +1541,7 @@ async def _resolve_object_permission_for_unresolvable_team( ) -async def _user_api_key_auth_builder( +async def user_api_key_auth_builder( request: Request, api_key: str, azure_api_key_header: str, @@ -1765,8 +1791,8 @@ async def _user_api_key_auth_builder( user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, parent_otel_span=parent_otel_span, - request_headers=_safe_get_request_headers(request), - request_method=RouteChecks._get_request_method(request=request), + request_headers=safe_get_request_headers(request), + request_method=RouteChecks.get_request_method(request=request), ) is_proxy_admin: Final = result["is_proxy_admin"] @@ -1862,9 +1888,11 @@ async def _user_api_key_auth_builder( ) skip_budget_checks = False if model is not None and llm_router is not None: - from litellm.proxy.auth.auth_checks import _is_model_cost_zero + from litellm.proxy.auth.auth_checks import is_model_cost_zero - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + skip_budget_checks = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_model_cost_zero(model=model, llm_router=llm_router) + ) if skip_budget_checks: verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) @@ -1950,7 +1978,7 @@ async def _user_api_key_auth_builder( _end_user_object = None end_user_params: Final = {} - raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request)) + raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, safe_get_request_headers(request)) end_user_id = await resolve_and_validate_end_user_id( raw_end_user_id=raw_end_user_id, prisma_client=prisma_client, @@ -2070,7 +2098,7 @@ async def _user_api_key_auth_builder( if expiry_time.tzinfo is None or expiry_time.tzinfo.utcoffset(expiry_time) is None: expiry_time = expiry_time.replace(tzinfo=timezone.utc) if expiry_time < current_time: - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hash_token(api_key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -2151,7 +2179,7 @@ async def _user_api_key_auth_builder( start_time=start_time, ) asyncio.create_task( - _cache_key_object( + cache_key_object( hashed_token=hash_token(master_key), user_api_key_obj=_user_api_key_obj, user_api_key_cache=user_api_key_cache, @@ -2248,7 +2276,7 @@ async def _user_api_key_auth_builder( _end_user_object=_end_user_object, ) except Exception as e: - return await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + return await UserAPIKeyAuthExceptionHandler.handle_authentication_error( e=e, request=request, request_data=request_data, @@ -2259,6 +2287,9 @@ async def _user_api_key_auth_builder( ) +_user_api_key_auth_builder: Final = user_api_key_auth_builder + + async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of existing shared authorization checks request: Request, request_data: dict[str, object], @@ -2303,7 +2334,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e ## base case ## key is disabled if valid_token.blocked is True: raise Exception("Key is blocked. Update via `/key/unblock` if you're an admin.") - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route=route, @@ -2359,9 +2390,11 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e ) skip_budget_checks = False if model is not None and llm_router is not None: - from litellm.proxy.auth.auth_checks import _is_model_cost_zero + from litellm.proxy.auth.auth_checks import is_model_cost_zero - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + skip_budget_checks = is_model_cost_zero( # rebind-ok: pre-existing rebinding on a rename-only line + model=model, llm_router=llm_router + ) if skip_budget_checks: verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) @@ -2412,7 +2445,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e if team_member_spend >= team_member_budget: # common_checks sends this alert on requests that get past here, so only the # request rejected here sends it from the builder. - _team_member_max_budget_alert_check( + team_member_max_budget_alert_check( team_id=_team_id, team_alias=valid_token.team_alias, team_metadata=valid_token.team_metadata, @@ -2461,7 +2494,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e # Check 4. Max Budget Alert Check (runs before budget enforcement # so multi-threshold 100% alerts fire on the request that crosses # max_budget, before BudgetExceededError is raised below) - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -2469,14 +2502,14 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e # Check 5. Token Spend is under budget if RouteChecks.is_llm_api_route(route=route): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, ) # Check 6. Soft Budget Check - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -2613,7 +2646,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY with tracer.trace("litellm.proxy.auth.get_global_proxy_spend"): - global_proxy_spend = await _fetch_global_spend_with_event_coordination( + global_proxy_spend = await fetch_global_spend_with_event_coordination( # rebind-ok: pre-existing rebinding on a rename-only line cache_key=cache_key, user_api_key_cache=user_api_key_cache, prisma_client=prisma_client, @@ -2803,7 +2836,7 @@ def is_no_auth_dev_mode(master_key: str | None, general_settings: Mapping[str, o @tracer.wrap() -async def _run_centralized_common_checks( +async def run_centralized_common_checks( user_api_key_auth_obj: UserAPIKeyAuth, request: Request, request_data: dict[str, object], @@ -2884,7 +2917,7 @@ async def _run_centralized_common_checks( key_end_user_budget_id: Final = get_key_end_user_budget_id(user_api_key_auth_obj.metadata) end_user_id = user_api_key_auth_obj.end_user_id if end_user_id is None: - raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request)) + raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, safe_get_request_headers(request)) end_user_id = await resolve_and_validate_end_user_id( raw_end_user_id=raw_end_user_id, prisma_client=prisma_client, @@ -3164,6 +3197,9 @@ async def _run_centralized_common_checks( release_spend_counter_batch() +_run_centralized_common_checks: Final = run_centralized_common_checks + + async def _noop_none() -> None: """Sentinel coroutine for asyncio.gather when a fetch is unnecessary (e.g. token has no team_id). Keeps the result tuple positional.""" @@ -3271,7 +3307,7 @@ def _should_skip_budget_checks( team_id=team_id, ) if model is not None and llm_router is not None: - return _is_model_cost_zero(model=model, llm_router=llm_router) + return is_model_cost_zero(model=model, llm_router=llm_router) return False @@ -3290,7 +3326,7 @@ def _resolve_request_principal(request: Request, valid_token: UserAPIKeyAuth) -> TrustedProxyConfig(use_forwarded_for=bool(cidrs), trusted_proxy_cidrs=cidrs), ) auth_method: Final = AuthMethod.BEARER_JWT if valid_token.jwt_claims else AuthMethod.API_KEY - return IdentityStore._principal_from_key( + return IdentityStore.principal_from_key( valid_token, auth_method=auth_method, network=network, @@ -3362,7 +3398,7 @@ async def _authorize_authenticated_request( billable=request_data.get("method") in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"), ) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=user_api_key_auth_obj, request=request, request_data=authorized_data, @@ -3370,7 +3406,7 @@ async def _authorize_authenticated_request( force_virtual_key_checks=force_virtual_key_checks, ) except Exception as e: - return await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + return await UserAPIKeyAuthExceptionHandler.handle_authentication_error( e=e, request=request, request_data=request_data, @@ -3393,7 +3429,7 @@ async def _authorize_authenticated_request( user_api_key_cache, ) - raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request)) + raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, safe_get_request_headers(request)) if raw_end_user_id is not None: resolved_end_user_id: Final = await resolve_and_validate_end_user_id( raw_end_user_id=raw_end_user_id, @@ -3475,7 +3511,7 @@ def _seed_request_destinations(user_api_key_dict: UserAPIKeyAuth, request: Reque set_request_destinations( deliverable_destinations( - resolve_tenant_otel_destinations(user_api_key_dict, _safe_get_request_headers(request)), + resolve_tenant_otel_destinations(user_api_key_dict, safe_get_request_headers(request)), fan_out_provider(), ) ) @@ -3518,7 +3554,7 @@ async def user_api_key_auth( spend_counter_batch_scope(_spend_counter_redis_cache()), ): try: - user_api_key_auth_obj: Final = await _user_api_key_auth_builder( + user_api_key_auth_obj: Final = await user_api_key_auth_builder( request=request, api_key=api_key, azure_api_key_header=azure_api_key_header, @@ -3537,7 +3573,7 @@ async def user_api_key_auth( raise user_api_key_auth_obj.budget_reservation = None user_api_key_auth_obj.agent_caller = agent_caller_from_headers( - _safe_get_request_headers(request), user_api_key_auth_obj + safe_get_request_headers(request), user_api_key_auth_obj ) _seed_request_destinations(user_api_key_auth_obj, request) @@ -3611,7 +3647,7 @@ async def _return_user_api_key_auth_obj( ) ) - retrieved_user_role: Final = user_role or _get_user_role(user_obj=user_obj) or LitellmUserRoles.INTERNAL_USER + retrieved_user_role: Final = user_role or get_user_role(user_obj=user_obj) or LitellmUserRoles.INTERNAL_USER user_api_key_kwargs: Final = { "api_key": api_key, @@ -3628,7 +3664,7 @@ async def _return_user_api_key_auth_obj( user_max_budget=getattr(user_obj, "max_budget", None), user_model_max_budget=getattr(user_obj, "model_max_budget", None), ) - if user_obj is not None and _is_user_proxy_admin(user_obj=user_obj): + if user_obj is not None and is_user_proxy_admin(user_obj=user_obj): user_api_key_kwargs.update( user_role=LitellmUserRoles.PROXY_ADMIN, ) @@ -3659,7 +3695,7 @@ def get_api_key_from_custom_header(request: Request, custom_litellm_key_header_n ) custom_api_key: Final = _headers.get(custom_litellm_key_header_name) if custom_api_key: - api_key = _get_bearer_token(api_key=custom_api_key) + api_key = get_bearer_token(api_key=custom_api_key) # rebind-ok: pre-existing rebinding on a rename-only line verbose_proxy_logger.debug( "Found custom API key using header: %s, setting api_key=%s", custom_litellm_key_header_name, @@ -3756,7 +3792,7 @@ async def _lookup_end_user_and_apply_budget( return valid_token, end_user_object -async def _enforce_key_and_fallback_model_access( +async def enforce_key_and_fallback_model_access( *, valid_token: UserAPIKeyAuth, request_data: dict, @@ -3816,6 +3852,9 @@ async def _enforce_key_and_fallback_model_access( ) +_enforce_key_and_fallback_model_access: Final = enforce_key_and_fallback_model_access + + async def _run_post_custom_auth_checks( valid_token: UserAPIKeyAuth, request: Request, @@ -3849,7 +3888,7 @@ async def _run_post_custom_auth_checks( # custom_auth_run_common_checks is set. Enforce it here on that path # so an over-budget end user can't keep making requests. if end_user_object is not None and not general_settings.get("custom_auth_run_common_checks", False): - await _check_end_user_budget(end_user_obj=end_user_object, route=route) + await check_end_user_budget(end_user_obj=end_user_object, route=route) # 2. Check token expiry if valid_token.expires is not None: @@ -3869,7 +3908,7 @@ async def _run_post_custom_auth_checks( ) if general_settings.get("custom_auth_run_common_checks", False): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route=route, @@ -3892,7 +3931,7 @@ async def _run_post_custom_auth_checks( # every budget check for these; this path did not, so the same request could # be refused under custom auth and served under the other two. skip_budget_checks: Final = ( - _is_model_cost_zero(model=current_model, llm_router=llm_router) + is_model_cost_zero(model=current_model, llm_router=llm_router) if current_model is not None and llm_router is not None else False ) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 2972b5043da..413d5a7057f 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -36,14 +36,17 @@ from litellm.proxy.common_request_processing import ( request_litellm_call_id, ) from litellm.proxy.common_utils.callback_utils import sanitize_openai_provider_metadata -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_headers, get_custom_llm_provider_from_request_query, ) -from litellm.proxy.openai_files_endpoints.common_utils import ( +from litellm.proxy.openai_files_endpoints.common_utils import ( # noqa: F401 # legacy module exports BATCH_CREATE_HIDDEN_PARAM, - _is_base64_encoded_unified_file_id, + _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export add_deployment_model_info, add_internal_model_credentials, apply_team_provider_credentials, @@ -59,6 +62,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( get_model_id_from_unified_batch_id, get_models_from_unified_file_id, get_original_file_id, + is_base64_encoded_unified_file_id, is_litellm_executed_batch, prepare_data_with_credentials, update_batch_in_database, @@ -269,7 +273,7 @@ async def create_batch( data: dict = {} try: - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line verbose_proxy_logger.debug( "Request received by LiteLLM:\n%s", json.dumps(data, indent=4), @@ -341,7 +345,9 @@ async def create_batch( model_from_file_id = None if input_file_id: model_from_file_id = decode_model_from_file_id(input_file_id) - unified_file_id = _is_base64_encoded_unified_file_id(input_file_id) + unified_file_id = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_base64_encoded_unified_file_id(input_file_id) + ) # SCENARIO 1: File ID is encoded with model info if model_from_file_id is not None and input_file_id: @@ -587,7 +593,7 @@ async def retrieve_batch( ) data = cast(dict, _retrieve_batch_request) - unified_batch_id: Final = _is_base64_encoded_unified_file_id(batch_id) + unified_batch_id: Final = is_base64_encoded_unified_file_id(batch_id) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) ( @@ -891,7 +897,7 @@ async def list_batches( ) # Include original request and headers in the data - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) ( data, @@ -1084,7 +1090,7 @@ async def cancel_batch( ) data = cast(dict, _cancel_batch_request) - unified_batch_id: Final = _is_base64_encoded_unified_file_id(batch_id) + unified_batch_id: Final = is_base64_encoded_unified_file_id(batch_id) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) ( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index fcc50e57fd5..317203cc2d5 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -28,6 +28,7 @@ from fastapi import HTTPException, Request, status from fastapi.responses import JSONResponse, Response, StreamingResponse from pydantic import BaseModel, TypeAdapter, ValidationError from starlette.types import Receive, Scope, Send +from typing_extensions import Never import litellm from litellm._logging import redact_internal_details_from_client_message, verbose_proxy_logger @@ -118,7 +119,11 @@ from litellm.proxy.native_compaction import with_proxy_compaction_executor from litellm.proxy.route_llm_request import ( route_request, ) -from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports + ProxyLogging, + _check_and_merge_model_level_guardrails, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + check_and_merge_model_level_guardrails, +) from litellm.router import Router from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.router_utils.common_utils import resolve_model_group_alias @@ -290,13 +295,16 @@ def resolve_litellm_call_id(client_call_id: str | None) -> str: return str(uuid.uuid4()) -def _should_return_raw_model_name(request_data: dict[str, object]) -> bool: +def should_return_raw_model_name(request_data: dict[str, object]) -> bool: return any( isinstance(metadata, dict) and metadata.get(RETURN_RAW_MODEL_NAME_METADATA_KEY) is True for metadata in (request_data.get("metadata"), request_data.get("litellm_metadata")) ) +_should_return_raw_model_name: Final = should_return_raw_model_name + + def _apply_client_disconnect_metadata(target_metadata: dict[str, object] | None) -> None: if target_metadata is None: return @@ -1302,7 +1310,7 @@ async def open_sse_before_first_byte( ) -def _is_azure_model_router_request(model: str, hidden_params: Mapping[str, object] | None = None) -> bool: +def is_azure_model_router_request(model: str, hidden_params: Mapping[str, object] | None = None) -> bool: """ Check if a request went down the Azure Model Router route. @@ -1323,6 +1331,9 @@ def _is_azure_model_router_request(model: str, hidden_params: Mapping[str, objec return AzureFoundryModelInfo.is_model_router_call(model=model, hidden_params=hidden_params) +_is_azure_model_router_request: Final = is_azure_model_router_request + + def _override_openai_response_model( *, response_obj: object, @@ -1388,7 +1399,7 @@ def _override_openai_response_model( return # Check if this is an Azure Model Router request - if so, preserve the actual model used - if _is_azure_model_router_request(requested_model, hidden_params): + if is_azure_model_router_request(requested_model, hidden_params): verbose_proxy_logger.debug( "%s: Azure Model Router detected - preserving actual model used from response instead of overriding to router model.", log_context, @@ -2074,9 +2085,9 @@ class ProxyBaseLLMRequestProcessing: # Store queue time in metadata after add_litellm_data_to_request to ensure it's preserved if queue_time_seconds is not None: - from litellm.proxy.litellm_pre_call_utils import _get_metadata_variable_name + from litellm.proxy.litellm_pre_call_utils import get_metadata_variable_name - _metadata_variable_name: Final = _get_metadata_variable_name(request) + _metadata_variable_name: Final = get_metadata_variable_name(request) if _metadata_variable_name not in self.data: self.data[_metadata_variable_name] = {} if not isinstance(self.data[_metadata_variable_name], dict): @@ -2200,11 +2211,11 @@ class ProxyBaseLLMRequestProcessing: merged_for_requested: Final = ( self.data if rate_limited_model is None - else _check_and_merge_model_level_guardrails( + else check_and_merge_model_level_guardrails( data=self.data, llm_router=llm_router, trust_client_model_info=False, model_alias=rate_limited_model ) ) - self.data = _check_and_merge_model_level_guardrails( + self.data = check_and_merge_model_level_guardrails( data=merged_for_requested, llm_router=llm_router, trust_client_model_info=False, @@ -2803,7 +2814,7 @@ class ProxyBaseLLMRequestProcessing: proxy_logging_obj=proxy_logging_obj, request=request, restamp_model=( - None if _should_return_raw_model_name(self.data) else requested_model_from_client + None if should_return_raw_model_name(self.data) else requested_model_from_client ), ) selected_data_generator = wrap_sse_stream_with_keepalive_pings( @@ -2942,7 +2953,7 @@ class ProxyBaseLLMRequestProcessing: response_obj=response, requested_model=requested_model_from_client, log_context=f"litellm_call_id={logging_obj.litellm_call_id}", - return_raw_model_name=_should_return_raw_model_name(self.data), + return_raw_model_name=should_return_raw_model_name(self.data), ) fastapi_response.headers.update( @@ -3235,9 +3246,9 @@ class ProxyBaseLLMRequestProcessing: because should_run_guardrail treats it as matching every hook. """ from litellm.proxy.proxy_server import llm_router - from litellm.proxy.utils import _check_and_merge_model_level_guardrails + from litellm.proxy.utils import check_and_merge_model_level_guardrails - guardrail_data: Final = _check_and_merge_model_level_guardrails(data=self.data, llm_router=llm_router) + guardrail_data: Final = check_and_merge_model_level_guardrails(data=self.data, llm_router=llm_router) for cb in litellm.callbacks: if not isinstance(cb, CustomGuardrail): continue @@ -3558,11 +3569,13 @@ class ProxyBaseLLMRequestProcessing: try: from litellm.proxy.proxy_server import llm_router as _global_llm_router from litellm.proxy.utils import ( - _check_and_merge_model_level_guardrails, + check_and_merge_model_level_guardrails, stream_gated_guardrail_names, ) - guardrail_data = _check_and_merge_model_level_guardrails(data=captured_data, llm_router=_global_llm_router) + guardrail_data: Final = check_and_merge_model_level_guardrails( + data=captured_data, llm_router=_global_llm_router + ) stream_gated: Final = stream_gated_guardrail_names(captured_data, captured_user_api_key_dict) for cb in litellm.callbacks: if not isinstance(cb, CustomGuardrail): @@ -3636,13 +3649,13 @@ class ProxyBaseLLMRequestProcessing: if isinstance(e, RouterRateLimitError) and e.cooldown_time > 0: headers["retry-after"] = str(math.ceil(e.cooldown_time)) - async def _handle_llm_api_exception( + async def handle_llm_api_exception( self, e: Exception, user_api_key_dict: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging, version: str | None = None, - ): + ) -> Never: """Raises ProxyException (OpenAI API compatible) if an exception is raised""" log_llm_api_exception(e, self.litellm_call_id) # Allow callbacks to transform the error response @@ -3787,6 +3800,8 @@ class ProxyBaseLLMRequestProcessing: headers=safe_headers, ) + _handle_llm_api_exception = handle_llm_api_exception + ######################################################### # Proxy Level Streaming Data Generator ######################################################### @@ -3820,7 +3835,7 @@ class ProxyBaseLLMRequestProcessing: return serialize @staticmethod - async def _finalize_streaming_generator_cleanup( + async def finalize_streaming_generator_cleanup( request: Request | None, request_data: dict, response: Any, @@ -3840,7 +3855,7 @@ class ProxyBaseLLMRequestProcessing: ) if recorded_client_disconnect: deferred_stream_logging_armed: Final = _deferred_stream_logging_is_armed(request_data) - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) # A disconnect-time success event (the deferred-guardrail flush # above, or the partial-spend billing below) releases the # request's max_parallel_requests slot through the limiter's @@ -3858,7 +3873,7 @@ class ProxyBaseLLMRequestProcessing: and proxy_logging_obj is not None and user_api_key_dict is not None ): - await proxy_logging_obj._arelease_max_parallel_requests_on_disconnect(user_api_key_dict) + await proxy_logging_obj.arelease_max_parallel_requests_on_disconnect(user_api_key_dict) if hasattr(response, "aclose"): try: @@ -3877,6 +3892,8 @@ class ProxyBaseLLMRequestProcessing: ): await logging_obj.invalidate_baseline_cache_estimate("incomplete_response", completed=True) + _finalize_streaming_generator_cleanup = finalize_streaming_generator_cleanup + @staticmethod async def async_streaming_data_generator( response: object, @@ -3913,7 +3930,7 @@ class ProxyBaseLLMRequestProcessing: # consumed, and cost injection is a no-op -- so the per-chunk coroutine # await, response-string materialization, and cost-injection call are # pure overhead on the streaming hot path (the default config). - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() cost_injection_enabled: Final = bool(getattr(litellm, "include_cost_in_streaming_usage", False)) fast_path = not caps.has_streaming_chunk_override and not caps.has_guardrail and not cost_injection_enabled debug_enabled: Final = verbose_proxy_logger.isEnabledFor(logging.DEBUG) @@ -3956,7 +3973,7 @@ class ProxyBaseLLMRequestProcessing: str_so_far += str(chunk.get("content", "")) model_name = request_data.get("model", "") - chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection( + chunk = ProxyBaseLLMRequestProcessing.process_chunk_with_cost_injection( chunk, model_name, request_data.get("litellm_logging_obj") ) @@ -4022,7 +4039,7 @@ class ProxyBaseLLMRequestProcessing: seal: Final = "" if seal_open_frame is None else seal_open_frame(recent_tail) yield seal + error_frame if seal else error_frame finally: - await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( + await ProxyBaseLLMRequestProcessing.finalize_streaming_generator_cleanup( request=request, request_data=request_data, response=response, @@ -4071,18 +4088,18 @@ class ProxyBaseLLMRequestProcessing: @overload @staticmethod - def _process_chunk_with_cost_injection( + def process_chunk_with_cost_injection( chunk: bytes, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None ) -> bytes: ... @overload @staticmethod - def _process_chunk_with_cost_injection( + def process_chunk_with_cost_injection( chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None ) -> object: ... @staticmethod - def _process_chunk_with_cost_injection( + def process_chunk_with_cost_injection( chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None ) -> object: """ @@ -4131,6 +4148,8 @@ class ProxyBaseLLMRequestProcessing: return chunk + _process_chunk_with_cost_injection = process_chunk_with_cost_injection + @staticmethod def _inject_cost_into_sse_frame_str( frame_str: str, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None diff --git a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py index 519e604783e..2d874f7b473 100644 --- a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py +++ b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py @@ -6,10 +6,12 @@ from typing import TYPE_CHECKING, Final from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger -from litellm.proxy.common_utils.config_sync_pubsub import ( - _ConfigSyncPubSub, - _pubsub_capable_client, +from litellm.proxy.common_utils.config_sync_pubsub import ( # noqa: F401 # legacy module exports + ConfigSyncPubSub, + _ConfigSyncPubSub, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _pubsub_capable_client, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export coordination_redis_cache, + pubsub_capable_client, ) from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET @@ -75,7 +77,7 @@ def _message_from_data(data: object) -> _CacheInvalidationMessage | None: async def _publish_to_redis(redis_cache: "RedisCache", cache_key: str, message: str) -> None: try: - client: Final = _pubsub_capable_client(redis_cache) + client: Final = pubsub_capable_client(redis_cache) async with _in_flight_publishes: await client.publish(auth_cache_invalidation_channel(redis_cache), message) except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors @@ -181,7 +183,7 @@ class AuthCacheInvalidationSubscriber: backoff_seconds = _BACKOFF_INITIAL_SECONDS # rebind-ok: exponential backoff accumulator across reconnects while True: try: - client = _pubsub_capable_client(self._redis_cache) + client = pubsub_capable_client(self._redis_cache) pubsub = client.pubsub() try: await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache)) @@ -200,7 +202,7 @@ class AuthCacheInvalidationSubscriber: await asyncio.sleep(backoff_seconds) backoff_seconds = min(backoff_seconds * 2, _BACKOFF_MAX_SECONDS) - async def _consume(self, pubsub: _ConfigSyncPubSub) -> None: + async def _consume(self, pubsub: ConfigSyncPubSub) -> None: while True: message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS) if message is None: @@ -222,7 +224,7 @@ class AuthCacheInvalidationSubscriber: additional_cache.delete_cache(parsed.cache_key) @staticmethod - async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None: + async def _close_pubsub(pubsub: ConfigSyncPubSub) -> None: try: await pubsub.aclose() except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection diff --git a/litellm/proxy/common_utils/cache_aware_routing.py b/litellm/proxy/common_utils/cache_aware_routing.py index d0685218c87..d5ae9aa9f66 100644 --- a/litellm/proxy/common_utils/cache_aware_routing.py +++ b/litellm/proxy/common_utils/cache_aware_routing.py @@ -243,7 +243,7 @@ async def _choose_cached_model( from litellm.proxy import proxy_server from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner + PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner ) from litellm.router_strategy.complexity_router.context_compaction import compaction_pending @@ -289,7 +289,7 @@ async def _choose_cached_model( if body is None: return None limiter: Final = proxy_server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter") - if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + if not isinstance(limiter, PROXY_MaxParallelRequestsHandler_v3): return None def counter_for_model(model_name: str) -> TokenCounter: diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 9ab39fa9303..4fa7709e5c7 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -181,7 +181,7 @@ def initialize_callbacks_on_proxy( imported_list.append(callback) elif isinstance(callback, str) and callback == "presidio": from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) presidio_logging_only: bool | None = litellm_settings.get("presidio_logging_only", None) @@ -196,7 +196,7 @@ def initialize_callbacks_on_proxy( "logging_only": presidio_logging_only, **_presidio_params, } - pii_masking_object = _OPTIONAL_PresidioPIIMasking(**params) + pii_masking_object = OPTIONAL_PresidioPIIMasking(**params) imported_list.append(pii_masking_object) elif isinstance(callback, str) and callback == "llamaguard_moderations": try: @@ -324,7 +324,7 @@ def initialize_callbacks_on_proxy( imported_list.append(banned_keywords_obj) elif isinstance(callback, str) and callback == "detect_prompt_injection": from litellm.proxy.hooks.prompt_injection_detection import ( - _OPTIONAL_PromptInjectionDetection, + OPTIONAL_PromptInjectionDetection, ) prompt_injection_params = None @@ -332,20 +332,20 @@ def initialize_callbacks_on_proxy( prompt_injection_params_in_config = litellm_settings["prompt_injection_params"] prompt_injection_params = LiteLLMPromptInjectionParams(**prompt_injection_params_in_config) - prompt_injection_detection_obj = _OPTIONAL_PromptInjectionDetection( + prompt_injection_detection_obj = OPTIONAL_PromptInjectionDetection( prompt_injection_params=prompt_injection_params, ) imported_list.append(prompt_injection_detection_obj) elif isinstance(callback, str) and callback == "batch_redis_requests": from litellm.proxy.hooks.batch_redis_get import ( - _PROXY_BatchRedisRequests, + PROXY_BatchRedisRequests, ) - batch_redis_obj = _PROXY_BatchRedisRequests() + batch_redis_obj = PROXY_BatchRedisRequests() imported_list.append(batch_redis_obj) elif isinstance(callback, str) and callback == "azure_content_safety": from litellm.proxy.hooks.azure_content_safety import ( - _PROXY_AzureContentSafety, + PROXY_AzureContentSafety, ) azure_content_safety_params = litellm_settings["azure_content_safety_params"] @@ -353,7 +353,7 @@ def initialize_callbacks_on_proxy( if v is not None and isinstance(v, str) and v.startswith("os.environ/"): azure_content_safety_params[k] = get_secret(v) - azure_content_safety_obj = _PROXY_AzureContentSafety( + azure_content_safety_obj = PROXY_AzureContentSafety( **azure_content_safety_params, ) imported_list.append(azure_content_safety_obj) diff --git a/litellm/proxy/common_utils/config_sync_pubsub.py b/litellm/proxy/common_utils/config_sync_pubsub.py index b4ebb5fa876..579db290190 100644 --- a/litellm/proxy/common_utils/config_sync_pubsub.py +++ b/litellm/proxy/common_utils/config_sync_pubsub.py @@ -21,10 +21,13 @@ class _ConfigSyncPubSub(Protocol): def aclose(self) -> Awaitable[object]: ... +ConfigSyncPubSub = _ConfigSyncPubSub + + class _ConfigSyncPubSubClient(Protocol): def publish(self, channel: str, message: str) -> Awaitable[int]: ... - def pubsub(self) -> _ConfigSyncPubSub: ... + def pubsub(self) -> ConfigSyncPubSub: ... CONFIG_SYNC_CHANNEL: Final = "litellm_proxy.config_change" @@ -82,13 +85,16 @@ def config_sync_channel(redis_cache: "RedisCache") -> str: return f"{redis_cache.namespace}:{CONFIG_SYNC_CHANNEL}" -def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient: +def pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient: return cast( # cast-ok: protocol view of the pub/sub-capable async redis client _ConfigSyncPubSubClient, redis_cache.init_pubsub_client(), # pyright: ignore[reportUnknownMemberType] # redis generics ) +_pubsub_capable_client: Final = pubsub_capable_client + + @dataclass(frozen=True, slots=True) class _ConfigChangeMessage: object_type: str @@ -102,7 +108,7 @@ async def publish_config_change(redis_cache: "RedisCache | None", object_type: s if redis_cache is None: return try: - client: Final = _pubsub_capable_client(redis_cache) + client: Final = pubsub_capable_client(redis_cache) await client.publish(config_sync_channel(redis_cache), _config_change_message_json(object_type)) except Exception as e: # noqa: BLE001 # best-effort publish; writes must never fail on redis errors verbose_proxy_logger.warning("config sync publish for %s failed: %s", object_type, e) @@ -222,7 +228,7 @@ class ConfigSyncSubscriber: backoff_seconds = self._backoff_initial_seconds while True: try: - client = _pubsub_capable_client(self._redis_cache) + client = pubsub_capable_client(self._redis_cache) pubsub = client.pubsub() try: await pubsub.subscribe(config_sync_channel(self._redis_cache)) @@ -241,7 +247,7 @@ class ConfigSyncSubscriber: await self._sleep(backoff_seconds) backoff_seconds = min(backoff_seconds * 2, self._backoff_max_seconds) - async def _consume(self, pubsub: _ConfigSyncPubSub) -> None: + async def _consume(self, pubsub: ConfigSyncPubSub) -> None: while True: message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS) if message is None: @@ -265,7 +271,7 @@ class ConfigSyncSubscriber: await self._sleep(seconds_until_next_resync) @staticmethod - async def _drain_pending(pubsub: _ConfigSyncPubSub) -> None: + async def _drain_pending(pubsub: ConfigSyncPubSub) -> None: while await pubsub.get_message(ignore_subscribe_messages=True, timeout=0) is not None: pass @@ -277,7 +283,7 @@ class ConfigSyncSubscriber: verbose_proxy_logger.warning("config sync resync callback failed: %s", e) @staticmethod - async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None: + async def _close_pubsub(pubsub: ConfigSyncPubSub) -> None: try: await pubsub.aclose() except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection diff --git a/litellm/proxy/common_utils/encrypt_decrypt_utils.py b/litellm/proxy/common_utils/encrypt_decrypt_utils.py index ae7240b8a7f..e5f6bc5dda8 100644 --- a/litellm/proxy/common_utils/encrypt_decrypt_utils.py +++ b/litellm/proxy/common_utils/encrypt_decrypt_utils.py @@ -12,18 +12,24 @@ from litellm._logging import verbose_proxy_logger # Legacy XSalsa20-Poly1305 (nacl) values carry no marker; the colon in the # prefix can never appear in base64url(nacl output), so the prefix check is an # unambiguous discriminator between the two formats on read. -_V2_GCM_PREFIX: Final = "v2:gcm:" +V2_GCM_PREFIX: Final = "v2:gcm:" + +_V2_GCM_PREFIX: Final = V2_GCM_PREFIX # general_settings key selecting the at-rest encryption algorithm for new writes. # Default preserves the legacy algorithm so existing deployments are byte-for-byte # unchanged until they explicitly opt in. Decrypt is always format-detecting, so # flipping this flag forward (or back) never strands previously-written data. -_ENCRYPTION_ALGORITHM_SETTING: Final = "encryption_algorithm" -_ALGO_AES_GCM: Final = "aes-256-gcm" +ENCRYPTION_ALGORITHM_SETTING: Final = "encryption_algorithm" + +_ENCRYPTION_ALGORITHM_SETTING: Final = ENCRYPTION_ALGORITHM_SETTING +ALGO_AES_GCM: Final = "aes-256-gcm" + +_ALGO_AES_GCM: Final = ALGO_AES_GCM _ALGO_XSALSA20: Final = "xsalsa20-poly1305" -def _get_salt_key(): +def get_salt_key() -> str | None: from litellm.proxy.proxy_server import master_key salt_key = os.getenv("LITELLM_SALT_KEY", None) @@ -34,6 +40,9 @@ def _get_salt_key(): return salt_key +_get_salt_key: Final = get_salt_key + + def _get_encryption_algorithm() -> str: """ Resolve the configured at-rest encryption algorithm for *new writes*. @@ -45,14 +54,14 @@ def _get_encryption_algorithm() -> str: try: from litellm.proxy.proxy_server import general_settings - algo: Final = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING, _ALGO_XSALSA20) + algo: Final = general_settings.get(ENCRYPTION_ALGORITHM_SETTING, _ALGO_XSALSA20) except Exception: # general_settings may not be importable in some contexts (e.g. SDK-only # use of these helpers). Fall back to the legacy algorithm. return _ALGO_XSALSA20 - if isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM: - return _ALGO_AES_GCM + if isinstance(algo, str) and algo.lower() == ALGO_AES_GCM: + return ALGO_AES_GCM return _ALGO_XSALSA20 @@ -92,18 +101,18 @@ def _open_aes_gcm(sealed: bytes, signing_key: str, aad: bytes | None) -> str: def _encrypt_aes_gcm(value: str, signing_key: str) -> str: """Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string.""" sealed: Final = _seal_aes_gcm(value=value, signing_key=signing_key, aad=None) - return _V2_GCM_PREFIX + base64.urlsafe_b64encode(sealed).decode("utf-8") + return V2_GCM_PREFIX + base64.urlsafe_b64encode(sealed).decode("utf-8") def _decrypt_aes_gcm(value: str, signing_key: str) -> str: """Decrypt a versioned ``v2:gcm:`` string produced by :func:`_encrypt_aes_gcm`.""" - sealed: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :]) + sealed: Final = base64.urlsafe_b64decode(value[len(V2_GCM_PREFIX) :]) return _open_aes_gcm(sealed=sealed, signing_key=signing_key, aad=None) def encrypt_bearer_token(value: str, prefix: str) -> str: """AES-256-GCM as unpadded base64url behind ``prefix``, which is also the AAD so a token can't change kind.""" - salt_key: Final = _get_salt_key() + salt_key: Final = get_salt_key() if not isinstance(salt_key, str): raise ValueError("Set LITELLM_SALT_KEY or a master key to mint bearer tokens") sealed: Final = _seal_aes_gcm(value=value, signing_key=salt_key, aad=prefix.encode("utf-8")) @@ -112,7 +121,7 @@ def encrypt_bearer_token(value: str, prefix: str) -> str: def decrypt_bearer_token(token: str, prefix: str) -> str | None: """None unless ``token`` came from :func:`encrypt_bearer_token` with the same ``prefix``.""" - salt_key: Final = _get_salt_key() + salt_key: Final = get_salt_key() if not isinstance(salt_key, str) or not token.startswith(prefix): return None encoded: Final = token.removeprefix(prefix) @@ -124,11 +133,11 @@ def decrypt_bearer_token(token: str, prefix: str) -> str | None: def encrypt_value_helper(value: str, new_encryption_key: str | None = None): - signing_key: Final = new_encryption_key or _get_salt_key() + signing_key: Final = new_encryption_key or get_salt_key() try: if isinstance(value, str): - if _get_encryption_algorithm() == _ALGO_AES_GCM: + if _get_encryption_algorithm() == ALGO_AES_GCM: # AES path: the v2:gcm: output is already a base64url string, so it # is returned directly with no extra base64 wrapper. return _encrypt_aes_gcm(value=value, signing_key=cast(str, signing_key)) @@ -160,7 +169,7 @@ def _legacy_ciphertext_bytes(value: str) -> bytes: def _decrypt_with_signing_key(value: str, signing_key: str) -> str: # Versioned AES-256-GCM values are detected before any base64 decode. # The prefix is the algorithm tag the legacy nacl format never carried. - if value.startswith(_V2_GCM_PREFIX): + if value.startswith(V2_GCM_PREFIX): return _decrypt_aes_gcm(value=value, signing_key=signing_key) return decrypt_value(value=_legacy_ciphertext_bytes(value), signing_key=signing_key) @@ -171,7 +180,7 @@ def decrypt_if_encrypted_with(value: str, signing_key: str) -> str | None: try: # base64 decoding skips characters outside its alphabet, so "" and "*" decode to no bytes, # which decrypt_value reads as an empty plaintext under any key. - decodes_to_nothing: Final = not value.startswith(_V2_GCM_PREFIX) and not _legacy_ciphertext_bytes(value) + decodes_to_nothing: Final = not value.startswith(V2_GCM_PREFIX) and not _legacy_ciphertext_bytes(value) return None if decodes_to_nothing else _decrypt_with_signing_key(value=value, signing_key=signing_key) except Exception: # noqa: BLE001 # base64, nacl and AES-GCM each raise their own "not a ciphertext" type return None @@ -183,7 +192,7 @@ def decrypt_value_helper( exception_type: Literal["debug", "error"] = "error", return_original_value: bool = False, ) -> str | None: - signing_key: Final = _get_salt_key() + signing_key: Final = get_salt_key() try: if isinstance(value, str): diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index fa288bad9fe..44360cec2bd 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -186,7 +186,7 @@ def is_otlp_trace_request(request: Request) -> bool: return request.method == "POST" and get_route_path(request.scope) in {"/v1/traces", "/v1/logs"} -async def _read_request_body(request: Request | None) -> dict: +async def read_request_body(request: Request | None) -> dict: """ Safely read the request body and parse it as JSON. @@ -208,7 +208,7 @@ async def _read_request_body(request: Request | None) -> dict: if _cached_request_body is not None: return _cached_request_body - _request_headers: Final[dict] = _safe_get_request_headers(request=request) + _request_headers: Final[dict] = safe_get_request_headers(request=request) content_type: Final = _request_headers.get("content-type", "") if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES: @@ -285,7 +285,7 @@ async def _read_request_body(request: Request | None) -> dict: ) # Cache the parsed result - _safe_set_request_parsed_body(request=request, parsed_body=parsed_body) + safe_set_request_parsed_body(request=request, parsed_body=parsed_body) return parsed_body except (json.JSONDecodeError, orjson.JSONDecodeError, ProxyException) as e: @@ -298,6 +298,9 @@ async def _read_request_body(request: Request | None) -> dict: return {} +_read_request_body: Final = read_request_body + + def is_opaque_audio_pass_through_request(route: str, content_type: str) -> bool: """Azure Speech bodies (raw audio, multipart uploads) are forwarded byte for byte, so auth must not consume them.""" media_type: Final = _normalize_media_type(content_type) @@ -309,7 +312,7 @@ def is_opaque_audio_pass_through_request(route: str, content_type: str) -> bool: async def read_raw_json_body(request: Request | None) -> bytes | None: if request is None or _safe_get_request_parsed_body(request=request) is None: return None - content_type: Final = _safe_get_request_headers(request=request).get("content-type", "") + content_type: Final = safe_get_request_headers(request=request).get("content-type", "") if _is_form_content_type(content_type): return None try: @@ -334,7 +337,7 @@ def get_client_requested_model(request: Request | None) -> str | None: return model if isinstance(model, str) else None -def _safe_get_request_query_params(request: Request | None) -> dict: +def safe_get_request_query_params(request: Request | None) -> dict: if request is None: return {} try: @@ -346,7 +349,10 @@ def _safe_get_request_query_params(request: Request | None) -> dict: return {} -def _safe_set_request_parsed_body( +_safe_get_request_query_params: Final = safe_get_request_query_params + + +def safe_set_request_parsed_body( request: Request | None, parsed_body: dict, ) -> None: @@ -358,6 +364,9 @@ def _safe_set_request_parsed_body( verbose_proxy_logger.debug("Unexpected error setting request parsed body - %s", e) +_safe_set_request_parsed_body: Final = safe_set_request_parsed_body + + def rewrite_request_model( request_data: dict[str, object], # mutable-ok: the request body is rewritten in place for every downstream reader request: Request | None, @@ -371,12 +380,12 @@ def rewrite_request_model( return cached_body: Final = _safe_get_request_parsed_body(request=request) body: Final = {**cached_body, "model": model} if cached_body is not None else request_data - _safe_set_request_parsed_body(request=request, parsed_body=body) - request._json = body - request._body = orjson.dumps(body) + safe_set_request_parsed_body(request=request, parsed_body=body) + request._json = body # pyright: ignore[reportPrivateUsage] # Starlette JSON cache + request._body = orjson.dumps(body) # pyright: ignore[reportPrivateUsage] # Starlette body cache -def _safe_get_request_headers(request: Request | None) -> dict: +def safe_get_request_headers(request: Request | None) -> dict: """ [Non-Blocking] Safely get the request headers. Caches the result on request.state to avoid re-creating dict(request.headers) per call. @@ -405,6 +414,9 @@ def _safe_get_request_headers(request: Request | None) -> dict: return headers +_safe_get_request_headers: Final = safe_get_request_headers + + def check_file_size_under_limit( request_data: dict, file: UploadFile, @@ -537,7 +549,7 @@ async def get_request_body(request: Request) -> dict[str, Any]: if request.method == "POST": content_type: Final = request.headers.get("content-type", "") if is_json_content_type(content_type): - return await _read_request_body(request) + return await read_request_body(request) elif _is_form_content_type(content_type): return await get_form_data(request) else: @@ -685,7 +697,7 @@ def populate_request_with_path_params(request_data: dict, request: Request) -> d dict: Updated request_data with path parameters and query parameters added """ # Add query parameters to request_data (for GET requests, etc.) - query_params: Final = _safe_get_request_query_params(request) + query_params: Final = safe_get_request_query_params(request) if query_params: for key, value in query_params.items(): # Don't overwrite existing values from request body diff --git a/litellm/proxy/common_utils/key_rotation_manager.py b/litellm/proxy/common_utils/key_rotation_manager.py index 352d024e20e..0ec196da6f5 100644 --- a/litellm/proxy/common_utils/key_rotation_manager.py +++ b/litellm/proxy/common_utils/key_rotation_manager.py @@ -20,8 +20,9 @@ from litellm.proxy._types import ( RegenerateKeyRequest, ) from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _calculate_key_rotation_time, +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _calculate_key_rotation_time, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + calculate_key_rotation_time, regenerate_key_fn, ) from litellm.proxy.utils import PrismaClient @@ -187,7 +188,7 @@ class KeyRotationManager: if isinstance(response, GenerateKeyResponse) and response.token_id and key.rotation_interval: # Calculate next rotation time using helper function now: Final = datetime.now(timezone.utc) - next_rotation_time: Final = _calculate_key_rotation_time(key.rotation_interval) + next_rotation_time: Final = calculate_key_rotation_time(key.rotation_interval) await VerificationTokenRepository(self.prisma_client).table.update( where={"token": response.token_id}, data={ diff --git a/litellm/proxy/common_utils/openai_endpoint_utils.py b/litellm/proxy/common_utils/openai_endpoint_utils.py index f85bbf3b380..ab44d32f368 100644 --- a/litellm/proxy/common_utils/openai_endpoint_utils.py +++ b/litellm/proxy/common_utils/openai_endpoint_utils.py @@ -7,7 +7,10 @@ from typing import Final from fastapi import Request from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) SENSITIVE_DATA_MASKER: Final = SensitiveDataMasker() @@ -55,7 +58,7 @@ async def get_custom_llm_provider_from_request_body(request: Request) -> str | N Safely reads the request body """ - request_body: Final[dict] = await _read_request_body(request=request) or {} + request_body: Final[dict] = await read_request_body(request=request) or {} if "custom_llm_provider" in request_body: return request_body["custom_llm_provider"] return None diff --git a/litellm/proxy/common_utils/openapi_schema_compat.py b/litellm/proxy/common_utils/openapi_schema_compat.py index 881b1fcc615..903e6c1cb24 100644 --- a/litellm/proxy/common_utils/openapi_schema_compat.py +++ b/litellm/proxy/common_utils/openapi_schema_compat.py @@ -40,7 +40,9 @@ def get_openapi_schema_with_compat( from pydantic_core import core_schema # Store original method - original_unknown_type_schema: Final = GenerateSchema._unknown_type_schema + original_unknown_type_schema: Final = ( + GenerateSchema._unknown_type_schema # pyright: ignore[reportPrivateUsage] # Pydantic schema internals + ) def patched_unknown_type_schema(self, obj): """Patch to handle openai.Timeout and other non-serializable types""" diff --git a/litellm/proxy/common_utils/rbac_utils.py b/litellm/proxy/common_utils/rbac_utils.py index 7c78d9a2470..b260fe127e4 100644 --- a/litellm/proxy/common_utils/rbac_utils.py +++ b/litellm/proxy/common_utils/rbac_utils.py @@ -51,10 +51,10 @@ async def check_feature_access_for_user( # Feature is disabled. Check if team/org admins are exempted. if general_settings.get(allow_team_admins_flag, False): from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_privileges, + user_has_admin_privileges, ) - is_admin: Final = await _user_has_admin_privileges( + is_admin: Final = await user_has_admin_privileges( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, diff --git a/litellm/proxy/common_utils/realtime_utils.py b/litellm/proxy/common_utils/realtime_utils.py index ff039754555..c641748843e 100644 --- a/litellm/proxy/common_utils/realtime_utils.py +++ b/litellm/proxy/common_utils/realtime_utils.py @@ -1,12 +1,16 @@ from functools import lru_cache +from typing import Final from litellm.constants import _REALTIME_BODY_CACHE_SIZE @lru_cache(maxsize=_REALTIME_BODY_CACHE_SIZE) -def _realtime_request_body(model: str | None) -> bytes: +def realtime_request_body(model: str | None) -> bytes: """ Generate the realtime websocket request body. Cached with LRU semantics to avoid repeated string formatting work while keeping memory usage bounded. """ return f'{{"model": "{model or ""}"}}'.encode() + + +_realtime_request_body: Final = realtime_request_body diff --git a/litellm/proxy/container_endpoints/endpoints.py b/litellm/proxy/container_endpoints/endpoints.py index 3142ea62b24..e11b0de3501 100644 --- a/litellm/proxy/container_endpoints/endpoints.py +++ b/litellm/proxy/container_endpoints/endpoints.py @@ -9,7 +9,10 @@ from litellm._logging import verbose_proxy_logger from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_body, get_custom_llm_provider_from_request_headers, @@ -89,7 +92,7 @@ async def create_container( ) # Read request body - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) # Extract custom_llm_provider using priority chain # Priority: headers > query params > request body > default @@ -125,7 +128,7 @@ async def create_container( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -245,7 +248,7 @@ async def list_containers( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -360,7 +363,7 @@ async def retrieve_container( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -466,7 +469,7 @@ async def delete_container( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/container_endpoints/handler_factory.py b/litellm/proxy/container_endpoints/handler_factory.py index 3989bdacef1..35bd69ef0bd 100644 --- a/litellm/proxy/container_endpoints/handler_factory.py +++ b/litellm/proxy/container_endpoints/handler_factory.py @@ -260,7 +260,7 @@ async def _process_binary_request( ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -349,7 +349,7 @@ async def _process_multipart_upload_request( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -438,7 +438,7 @@ async def _process_request( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/custom_hooks/custom_ui_sso_hook.py b/litellm/proxy/custom_hooks/custom_ui_sso_hook.py index 53002680829..9220022a648 100644 --- a/litellm/proxy/custom_hooks/custom_ui_sso_hook.py +++ b/litellm/proxy/custom_hooks/custom_ui_sso_hook.py @@ -5,7 +5,10 @@ from fastapi_sso.sso.base import OpenID from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) class CustomSSOLoginHandler(CustomLogger): @@ -22,7 +25,7 @@ class CustomSSOLoginHandler(CustomLogger): self, request: Request, ) -> OpenID: - request_headers_dict: Final = _safe_get_request_headers(request) + request_headers_dict: Final = safe_get_request_headers(request) verbose_logger.debug("inside custom ui sso sign in hook...") return OpenID( id=request_headers_dict.get("x-litellm-user-id") or "123", diff --git a/litellm/proxy/db/db_span.py b/litellm/proxy/db/db_span.py index 0cbd7c8db10..06b1a6aa11b 100644 --- a/litellm/proxy/db/db_span.py +++ b/litellm/proxy/db/db_span.py @@ -22,7 +22,12 @@ from typing import Final, TypeVar from litellm._logging import verbose_proxy_logger from litellm._service_logger import ServiceLogging, ServiceTypes -from litellm.proxy.db.log_db_metrics import _is_exception_related_to_db, claim_db_io, db_io_claimed +from litellm.proxy.db.log_db_metrics import ( # noqa: F401 # legacy module exports + _is_exception_related_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + claim_db_io, + db_io_claimed, + is_exception_related_to_db, +) _T = TypeVar("_T") @@ -75,7 +80,7 @@ async def db_span(call_type: str, table: str | None, operation: str | None = Non try: yield except Exception as e: - if service_logging is not None and _is_exception_related_to_db(e): + if service_logging is not None and is_exception_related_to_db(e): await _emit_failure(service_logging, call_type, event_metadata, start_time, e) raise if service_logging is None or not witness.touched: diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 7bbb0e63202..158a6c1a3c4 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1546,7 +1546,7 @@ class DBSpendUpdateWriter: else: - Regular flow of this method """ - if RedisUpdateBuffer._should_commit_spend_updates_to_redis(): + if RedisUpdateBuffer.should_commit_spend_updates_to_redis(): await self._commit_spend_updates_to_db_with_redis( prisma_client=prisma_client, n_retry_times=n_retry_times, @@ -1885,12 +1885,12 @@ class DBSpendUpdateWriter: ################## Tool Registry Upserts ################## await self._flush_tool_discovery_queue(prisma_client=prisma_client) - async def _commit_daily_tag_spend_to_db( + async def commit_daily_tag_spend_to_db( self, prisma_client: PrismaClient, n_retry_times: int, proxy_logging_obj: ProxyLogging, - ): + ) -> None: """ Commit only tag spend updates to database. This is called by a separate scheduler job at a longer interval. @@ -1904,12 +1904,14 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, ) - async def _commit_daily_tag_spend_to_db_with_redis( + _commit_daily_tag_spend_to_db = commit_daily_tag_spend_to_db + + async def commit_daily_tag_spend_to_db_with_redis( self, prisma_client: PrismaClient, n_retry_times: int, proxy_logging_obj: ProxyLogging, - ): + ) -> None: """ Commit daily tag spend updates using Redis buffering. @@ -1942,6 +1944,8 @@ class DBSpendUpdateWriter: cronjob_id=DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME, ) + _commit_daily_tag_spend_to_db_with_redis = commit_daily_tag_spend_to_db_with_redis + @staticmethod async def _commit_window_spend_updates( prisma_client: PrismaClient, @@ -2027,7 +2031,7 @@ class DBSpendUpdateWriter: verbose_proxy_logger.debug("_flush_tool_discovery_queue error (non-blocking): %s", e) @staticmethod - async def _handle_spend_update_failure( + async def handle_spend_update_failure( e: Exception, attempt: int, n_retry_times: int, @@ -2038,7 +2042,7 @@ class DBSpendUpdateWriter: ``lock_timeout`` (55P03), else re-raise. All three roll the transaction back before any increment applied, so re-sending the same batch cannot double-count.""" from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler - from litellm.proxy.utils import _raise_failed_update_spend_exception + from litellm.proxy.utils import raise_failed_update_spend_exception is_retryable = ( isinstance(e, DB_RETRY_SAFE_ERROR_TYPES) @@ -2046,7 +2050,7 @@ class DBSpendUpdateWriter: or PrismaDBExceptionHandler.is_lock_timeout_error(e) ) if not is_retryable or attempt >= n_retry_times: - _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) + raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) verbose_proxy_logger.warning( "Retrying spend update after retryable DB error (attempt %s/%s): %s", attempt + 1, @@ -2055,6 +2059,8 @@ class DBSpendUpdateWriter: ) await asyncio.sleep(random.uniform(2**attempt, 2 ** (attempt + 1))) + _handle_spend_update_failure = handle_spend_update_failure + async def _commit_spend_updates_to_db( self, prisma_client: PrismaClient, @@ -2088,7 +2094,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2132,7 +2138,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2162,7 +2168,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2192,7 +2198,7 @@ class DBSpendUpdateWriter: # Transaction succeeded, break out of retry loop break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2234,7 +2240,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2262,7 +2268,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await self._handle_spend_update_failure( + await self.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2389,7 +2395,7 @@ class DBSpendUpdateWriter: ) break except Exception as e: - await DBSpendUpdateWriter._handle_spend_update_failure( + await DBSpendUpdateWriter.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -2489,7 +2495,7 @@ class DBSpendUpdateWriter: """ Generic function to update daily spend for any entity type (user, team, org, tag, end_user, agent) """ - from litellm.proxy.utils import _raise_failed_update_spend_exception + from litellm.proxy.utils import raise_failed_update_spend_exception verbose_proxy_logger.debug( "Daily %s Spend transactions: %s", entity_type.capitalize(), len(daily_spend_transactions) @@ -2589,7 +2595,7 @@ class DBSpendUpdateWriter: if not is_retryable: raise if i >= n_retry_times: - _raise_failed_update_spend_exception( + raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj, @@ -2604,7 +2610,7 @@ class DBSpendUpdateWriter: ) except Exception as e: - _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) + raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) @staticmethod async def update_daily_user_spend( diff --git a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py index 3d24b0a0612..898aee493df 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -147,7 +147,7 @@ class RedisUpdateBuffer: self.redis_cache = redis_cache @staticmethod - def _should_commit_spend_updates_to_redis() -> bool: + def should_commit_spend_updates_to_redis() -> bool: """ Checks if the Pod should commit spend updates to Redis @@ -163,6 +163,8 @@ class RedisUpdateBuffer: return False return _use_redis_transaction_buffer + _should_commit_spend_updates_to_redis = should_commit_spend_updates_to_redis + @with_service_target(SPEND_QUEUE_TARGET) async def _store_transactions_in_redis( self, @@ -556,7 +558,7 @@ class RedisUpdateBuffer: max_rows: int = REDIS_SPEND_LOGS_BUFFER_MAX_ROWS, ) -> bool: """Park spend-log rows in Redis so they outlive this pod, dropping the oldest past ``max_rows``.""" - if self.redis_cache is None or len(rows) == 0 or not self._should_commit_spend_updates_to_redis(): + if self.redis_cache is None or len(rows) == 0 or not self.should_commit_spend_updates_to_redis(): return False try: buffer_size: Final = await self.redis_cache.async_rpush_and_trim( @@ -582,7 +584,7 @@ class RedisUpdateBuffer: @with_service_target(SPEND_QUEUE_TARGET) async def get_spend_logs_from_redis_buffer(self, limit: int) -> tuple[dict[str, object], ...]: """Atomically take up to ``limit`` parked spend-log rows out of Redis.""" - if self.redis_cache is None or not self._should_commit_spend_updates_to_redis(): + if self.redis_cache is None or not self.should_commit_spend_updates_to_redis(): return () popped: Final[str | list[str] | None] = await self.redis_cache.async_lpop( key=REDIS_SPEND_LOGS_BUFFER_KEY, diff --git a/litellm/proxy/db/log_db_metrics.py b/litellm/proxy/db/log_db_metrics.py index 988f93fb962..027dbe8e9fb 100644 --- a/litellm/proxy/db/log_db_metrics.py +++ b/litellm/proxy/db/log_db_metrics.py @@ -175,7 +175,7 @@ def log_db_metrics(func): return wrapper -def _is_exception_related_to_db(e: Exception) -> bool: +def is_exception_related_to_db(e: Exception) -> bool: """ Returns True if the exception is related to the DB """ @@ -186,6 +186,9 @@ def _is_exception_related_to_db(e: Exception) -> bool: return isinstance(e, (PrismaError, httpx.TransportError)) +_is_exception_related_to_db: Final = is_exception_related_to_db + + async def _handle_logging_db_exception( e: Exception, func: Callable, @@ -198,7 +201,7 @@ async def _handle_logging_db_exception( from litellm.proxy.proxy_server import proxy_logging_obj # don't log this as a DB Service Failure, if the DB did not raise an exception - if _is_exception_related_to_db(e) is not True: + if is_exception_related_to_db(e) is not True: return False try: diff --git a/litellm/proxy/decisions_endpoints/endpoints.py b/litellm/proxy/decisions_endpoints/endpoints.py index 7dba64ee791..e05d37963cb 100644 --- a/litellm/proxy/decisions_endpoints/endpoints.py +++ b/litellm/proxy/decisions_endpoints/endpoints.py @@ -87,14 +87,14 @@ async def decisions( model=str(data.get("model", "")), llm_provider="", ) - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=bad_request_error, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, version=version, ) except Exception as error: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=error, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/fine_tuning_endpoints/endpoints.py b/litellm/proxy/fine_tuning_endpoints/endpoints.py index 886b8da1454..d015112cf3a 100644 --- a/litellm/proxy/fine_tuning_endpoints/endpoints.py +++ b/litellm/proxy/fine_tuning_endpoints/endpoints.py @@ -16,8 +16,9 @@ from litellm.litellm_core_utils.hidden_params import set_hidden_param from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, +from litellm.proxy.openai_files_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_base64_encoded_unified_file_id, validate_managed_id_requirement, ) from litellm.proxy.utils import handle_exception_on_proxy @@ -150,7 +151,9 @@ async def create_fine_tuning_job( ) response: LiteLLMFineTuningJob | None = None if training_file: - unified_file_id = _is_base64_encoded_unified_file_id(training_file) + unified_file_id = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_base64_encoded_unified_file_id(training_file) + ) ## IF SO, Route based on that if unified_file_id: """ """ @@ -292,7 +295,9 @@ async def retrieve_fine_tuning_job( unified_finetuning_job_id: str | Literal[False] = False response: LiteLLMFineTuningJob | None = None if fine_tuning_job_id: - unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id) + unified_finetuning_job_id = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_base64_encoded_unified_file_id(fine_tuning_job_id) + ) if unified_finetuning_job_id: if llm_router is None: raise HTTPException( @@ -565,7 +570,9 @@ async def cancel_fine_tuning_job( unified_finetuning_job_id: str | Literal[False] = False response: LiteLLMFineTuningJob | None = None if fine_tuning_job_id: - unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id) + unified_finetuning_job_id = ( # rebind-ok: pre-existing rebinding on a rename-only line + is_base64_encoded_unified_file_id(fine_tuning_job_id) + ) if unified_finetuning_job_id: if llm_router is None: raise HTTPException( diff --git a/litellm/proxy/google_endpoints/agents_endpoints.py b/litellm/proxy/google_endpoints/agents_endpoints.py index 59ef45c7817..00b23d846c2 100644 --- a/litellm/proxy/google_endpoints/agents_endpoints.py +++ b/litellm/proxy/google_endpoints/agents_endpoints.py @@ -23,9 +23,11 @@ from fastapi.responses import ORJSONResponse from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_query_params, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_query_params, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, + safe_get_request_query_params, ) router: Final = APIRouter(tags=["gemini managed agents"]) @@ -92,7 +94,7 @@ def _merge_query_params_into_data(data: dict, request: Request) -> dict: headers. Use the ``litellm_params_template`` JSON body field on POST requests, or the JSON-encoded query parameter above for GET/DELETE. """ - query_params: Final = _safe_get_request_query_params(request) + query_params: Final = safe_get_request_query_params(request) if not query_params: return data @@ -172,7 +174,7 @@ async def create_gemini_agent( ``` """ srv: Final = _proxy_server_imports() - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) # Merge litellm_params_template (e.g. custom_llm_provider, api_key) into the request litellm_params_template: Final = data.pop("litellm_params_template", None) or {} if isinstance(litellm_params_template, dict): @@ -203,7 +205,7 @@ async def create_gemini_agent( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], @@ -260,7 +262,7 @@ async def list_gemini_agents( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], @@ -318,7 +320,7 @@ async def get_gemini_agent( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], @@ -376,7 +378,7 @@ async def delete_gemini_agent( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], @@ -434,7 +436,7 @@ async def list_gemini_agent_versions( version=srv["version"], ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=srv["proxy_logging_obj"], diff --git a/litellm/proxy/google_endpoints/endpoints.py b/litellm/proxy/google_endpoints/endpoints.py index dd5dc66d82b..87d2ee07818 100644 --- a/litellm/proxy/google_endpoints/endpoints.py +++ b/litellm/proxy/google_endpoints/endpoints.py @@ -6,7 +6,10 @@ from fastapi.responses import ORJSONResponse from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.types.llms.vertex_ai import TokenCountDetailsResponse router: Final = APIRouter( @@ -42,7 +45,7 @@ async def google_generate_content( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) if "model" not in data: data["model"] = model_name @@ -67,7 +70,7 @@ async def google_generate_content( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -103,7 +106,7 @@ async def google_stream_generate_content( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) if "model" not in data: data["model"] = model_name data["stream"] = True @@ -132,7 +135,7 @@ async def google_stream_generate_content( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -166,10 +169,10 @@ async def google_count_tokens(request: Request, model_name: str): ``` """ from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body from litellm.proxy.proxy_server import token_counter as internal_token_counter - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) contents: Final = data.get("contents", []) # Create TokenCountRequest for the internal endpoint from litellm.proxy._types import TokenCountRequest @@ -268,7 +271,7 @@ async def create_interaction( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) # Default to gemini provider for interactions if "custom_llm_provider" not in data: @@ -295,7 +298,7 @@ async def create_interaction( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -363,7 +366,7 @@ async def get_interaction( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -431,7 +434,7 @@ async def delete_interaction( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -499,7 +502,7 @@ async def cancel_interaction( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 1c9639fa9dc..8e1209401ab 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -43,7 +43,10 @@ from litellm.proxy.guardrails.guardrail_registry import ( parse_tolerant_litellm_params, ) from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, +) from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import GuardrailsRepository from litellm.types.guardrails import ( @@ -247,7 +250,7 @@ async def list_guardrails_v2( from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER from litellm.proxy.proxy_server import prisma_client - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) try: guardrails = ( @@ -942,7 +945,7 @@ async def list_guardrail_submissions( # Admin Viewer follows the read-parity rule: see all submissions like a # Proxy Admin would (no writes — registration / approval still gated # elsewhere by their own per-action checks). - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) visible_team_ids: list[str] | None = None if not is_admin: visible_team_ids = await _get_user_team_ids(user_api_key_dict) @@ -1021,7 +1024,7 @@ async def get_guardrail_submission( if prisma_client is None: raise HTTPException(status_code=500, detail="Prisma client not initialized") - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) try: row: Final = await _guardrails_table(prisma_client).find_unique(where={"guardrail_id": guardrail_id}) diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py index b69028594b2..5f57969c310 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/base.py @@ -26,7 +26,9 @@ AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH: Final = 1000 AZURE_CONTENT_SAFETY_DEFAULT_API_VERSION: Final = "2024-09-01" JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES: Final = "v1" -_RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses}) +RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses}) + +_RESPONSES_API_CALL_TYPES: Final = RESPONSES_API_CALL_TYPES def resolve_content_safety_api_version(configured: str | None) -> str: @@ -155,7 +157,7 @@ class AzureGuardrailBase: return get_last_user_message(messages) def get_user_prompt_from_request(self, data: Mapping[str, object], call_type: CallTypesLiteral) -> str | None: - if call_type in _RESPONSES_API_CALL_TYPES: + if call_type in RESPONSES_API_CALL_TYPES: responses_input: Final = data.get("input") if not isinstance(responses_input, (str, list)): return None diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py index 6a9c5aa4fbb..3967ef7152c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py @@ -27,7 +27,12 @@ from litellm.types.utils import ( GuardrailTracingDetail, ) -from .base import _RESPONSES_API_CALL_TYPES, AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH, AzureGuardrailBase +from .base import ( # noqa: F401 # legacy module exports + _RESPONSES_API_CALL_TYPES, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH, + RESPONSES_API_CALL_TYPES, + AzureGuardrailBase, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -249,7 +254,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai "Azure Prompt Shield: Running pre-call prompt scan, on call_type: %s", call_type, ) - if call_type not in _RESPONSES_API_CALL_TYPES and data.get("messages") is None: + if call_type not in RESPONSES_API_CALL_TYPES and data.get("messages") is None: verbose_proxy_logger.warning("Azure Prompt Shield: not running guardrail. No messages in data") return data user_prompt: Final = self.get_user_prompt_from_request(data, call_type) diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py index d9147cfb62b..27f36f2724b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py @@ -16,7 +16,11 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import CallTypesLiteral, GenericGuardrailAPIInputs, LLMResponseTypes -from .base import _RESPONSES_API_CALL_TYPES, AzureGuardrailBase +from .base import ( # noqa: F401 # legacy module exports + _RESPONSES_API_CALL_TYPES, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + RESPONSES_API_CALL_TYPES, + AzureGuardrailBase, +) if TYPE_CHECKING: from litellm.caching.caching import DualCache @@ -231,7 +235,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr "Azure Text Moderation: Running pre-call prompt scan, on call_type: %s", call_type, ) - if call_type not in _RESPONSES_API_CALL_TYPES and data.get("messages") is None: + if call_type not in RESPONSES_API_CALL_TYPES and data.get("messages") is None: verbose_proxy_logger.warning("Azure Text Moderation: not running guardrail. No messages in data") return data user_prompt: Final = self.get_user_prompt_from_request(data, call_type) diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py index 1e24b2aecbc..ab43d32e260 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py @@ -52,7 +52,10 @@ from litellm.types.utils import ( TextCompletionResponse, ) -from .cisco_ai_defense_mcp import _CiscoAIDefenseMcpMixin +from .cisco_ai_defense_mcp import ( # noqa: F401 # legacy module exports + CiscoAIDefenseMcpMixin, + _CiscoAIDefenseMcpMixin, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) if TYPE_CHECKING: from litellm.types.proxy.guardrails.guardrail_hooks.base import ( @@ -116,7 +119,7 @@ class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" -class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): +class CiscoAIDefenseGuardrail(CiscoAIDefenseMcpMixin, CustomGuardrail): """ Cisco AI Defense guardrail integration. diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py index 67ef05fc324..d7cc61ab511 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py @@ -218,7 +218,7 @@ class _CiscoAIDefenseMcpMixin: inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None) if inner is not None: - if _CiscoAIDefenseMcpMixin._replace_mcp_tool_response(inner, replacement_obj): + if CiscoAIDefenseMcpMixin._replace_mcp_tool_response(inner, replacement_obj): return True try: setattr(response_obj, "mcp_tool_call_response", replacement) @@ -229,7 +229,7 @@ class _CiscoAIDefenseMcpMixin: content: Final = getattr(response_obj, "content", None) if isinstance(content, list): content[:] = replacement - structured_replacement: Final = _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement) + structured_replacement: Final = CiscoAIDefenseMcpMixin._replacement_structured_content(replacement) if hasattr(response_obj, "structured_content"): try: setattr(response_obj, "structured_content", structured_replacement) @@ -250,12 +250,12 @@ class _CiscoAIDefenseMcpMixin: result: Final = response_obj.get("result") if isinstance(result, dict): result["content"] = replacement - result["structuredContent"] = _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement) + result["structuredContent"] = CiscoAIDefenseMcpMixin._replacement_structured_content(replacement) result["isError"] = True return True response_obj["result"] = { "content": replacement, - "structuredContent": _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement), + "structuredContent": CiscoAIDefenseMcpMixin._replacement_structured_content(replacement), "isError": True, } return True @@ -472,7 +472,7 @@ class _CiscoAIDefenseMcpMixin: return { "jsonrpc": "2.0", "id": response.get("id") or "litellm-mcp", - "result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response), + "result": CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response), } if isinstance(response, list): if response and all( @@ -484,7 +484,7 @@ class _CiscoAIDefenseMcpMixin: return { "jsonrpc": "2.0", "id": "litellm-mcp", - "result": _CiscoAIDefenseMcpMixin._build_mcp_result( + "result": CiscoAIDefenseMcpMixin._build_mcp_result( content=inner_content, source=response_fields ), } @@ -493,7 +493,7 @@ class _CiscoAIDefenseMcpMixin: return { "jsonrpc": "2.0", "id": "litellm-mcp", - "result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=response), + "result": CiscoAIDefenseMcpMixin._build_mcp_result(content=response), } model_dump: Final = getattr(response, "model_dump", None) if callable(model_dump): @@ -502,13 +502,13 @@ class _CiscoAIDefenseMcpMixin: except TypeError: dumped = model_dump() if isinstance(dumped, dict): - return _CiscoAIDefenseMcpMixin._normalize_mcp_response(dumped) + return CiscoAIDefenseMcpMixin._normalize_mcp_response(dumped) content = getattr(response, "content", None) if isinstance(content, list): return { "jsonrpc": "2.0", "id": "litellm-mcp", - "result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response), + "result": CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response), } return None @@ -536,9 +536,9 @@ class _CiscoAIDefenseMcpMixin: inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None) if inner is not None: - return _CiscoAIDefenseMcpMixin._set_mcp_tool_response_text(inner, text) + return CiscoAIDefenseMcpMixin._set_mcp_tool_response_text(inner, text) - content_list: Final = _CiscoAIDefenseMcpMixin._coerce_to_content_list(response_obj) + content_list: Final = CiscoAIDefenseMcpMixin._coerce_to_content_list(response_obj) replaced = False if isinstance(content_list, list): @@ -586,7 +586,7 @@ class _CiscoAIDefenseMcpMixin: return None inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None) if inner is not None: - return _CiscoAIDefenseMcpMixin._coerce_to_content_list(inner) + return CiscoAIDefenseMcpMixin._coerce_to_content_list(inner) content: Final = getattr(response_obj, "content", None) if isinstance(content, list): return content @@ -643,3 +643,6 @@ class _CiscoAIDefenseMcpMixin: if isinstance(direct, dict) and direct: return dict(direct) return None + + +CiscoAIDefenseMcpMixin = _CiscoAIDefenseMcpMixin diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py index 96365b24410..5e2b85250d2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/airline.py @@ -13,11 +13,14 @@ import json from pathlib import Path from typing import Any, Final -from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base import ( +from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base import ( # noqa: F401 # legacy module exports BaseCompetitorIntentChecker, - _compile_marker, - _count_signals, - _word_boundary_match, + _compile_marker, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _count_signals, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _word_boundary_match, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + compile_marker, + count_signals, + word_boundary_match, ) # Location/travel context: prepositions, travel verbs, booking nouns, entry/geo nouns. @@ -170,8 +173,8 @@ class AirlineCompetitorIntentChecker(BaseCompetitorIntentChecker): self._other_meaning_signals = list(merged.get("other_meaning_signals") or []) self._competitor_signals = list(merged.get("competitor_signals") or []) self._other_meaning_anchors = list(merged.get("other_meaning_anchors") or []) - self._explicit_competitor_marker = _compile_marker(merged.get("explicit_competitor_marker")) - self._explicit_other_meaning_marker = _compile_marker(merged.get("explicit_other_meaning_marker")) + self._explicit_competitor_marker = compile_marker(merged.get("explicit_competitor_marker")) + self._explicit_other_meaning_marker = compile_marker(merged.get("explicit_other_meaning_marker")) def _classify_ambiguous(self, text: str, token: str) -> tuple[str, float]: """Other meaning vs competitor using airline signals and explicit markers.""" @@ -179,21 +182,25 @@ class AirlineCompetitorIntentChecker(BaseCompetitorIntentChecker): if ( self._explicit_competitor_marker and self._explicit_competitor_marker.search(text_lower) - and _word_boundary_match(text_lower, token.lower()) + and word_boundary_match(text_lower, token.lower()) ): return "COMPETITOR", 0.85 if self._explicit_other_meaning_marker and self._explicit_other_meaning_marker.search(text_lower): return "OTHER_MEANING", 0.85 # Operational-only: baggage/lounge/check-in/refund with no comparison → product query - has_comparison: Final = _count_signals(text_lower, AIRLINE_COMPARISON_SIGNALS) > 0 - operational_count: Final = _count_signals(text_lower, AIRLINE_OPERATIONAL_SIGNALS) + has_comparison: Final = count_signals(text_lower, AIRLINE_COMPARISON_SIGNALS) > 0 + operational_count: Final = count_signals(text_lower, AIRLINE_OPERATIONAL_SIGNALS) if not has_comparison and operational_count > 0: return "OTHER_MEANING", 0.85 # Score: location/travel context vs airline context (no place-name list) - other_count = _count_signals(text_lower, self._other_meaning_signals) + other_count = count_signals( # rebind-ok: pre-existing rebinding on a rename-only line + text_lower, self._other_meaning_signals + ) if self._other_meaning_anchors: - other_count += _count_signals(text_lower, self._other_meaning_anchors) - comp_count: Final = _count_signals(text_lower, self._competitor_signals) + other_count += count_signals( # rebind-ok: pre-existing rebinding on a rename-only line + text_lower, self._other_meaning_anchors + ) + comp_count: Final = count_signals(text_lower, self._competitor_signals) total: Final = other_count + comp_count if total == 0: return "OTHER_MEANING", 0.5 diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py index 246e56441e5..eb801a412e0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/competitor_intent/base.py @@ -31,17 +31,23 @@ def normalize(text: str) -> str: return re.sub(r"\s+", " ", t) -def _word_boundary_match(text: str, token: str) -> bool: +def word_boundary_match(text: str, token: str) -> bool: """True if token appears as a word in text.""" return bool(re.search(r"\b" + re.escape(token) + r"\b", text)) -def _count_signals(text: str, patterns: list[str]) -> int: +_word_boundary_match: Final = word_boundary_match + + +def count_signals(text: str, patterns: list[str]) -> int: """Count how many of the patterns appear in text.""" return sum(1 for p in patterns if re.search(p, text, re.IGNORECASE)) -def _compile_marker(pattern: str | None) -> Pattern[str] | None: +_count_signals: Final = count_signals + + +def compile_marker(pattern: str | None) -> Pattern[str] | None: """Compile optional regex string to a pattern.""" if not pattern or not pattern.strip(): return None @@ -51,6 +57,9 @@ def _compile_marker(pattern: str | None) -> Pattern[str] | None: return None +_compile_marker: Final = compile_marker + + def text_for_entity_matching(text: str) -> str: """Letters-only variant for entity matching (e.g. split punctuation).""" t: Final = re.sub(r"[^\w\s]", " ", text) @@ -117,7 +126,7 @@ class BaseCompetitorIntentChecker: found: Final[list[tuple[str, str, bool]]] = [] seen: Final[set[tuple[str, str]]] = set() for token in self._competitor_tokens: - if not _word_boundary_match(normalized, token): + if not word_boundary_match(normalized, token): continue canonical = self.competitor_canonical.get(token, token) key = (token, canonical) @@ -139,7 +148,7 @@ class BaseCompetitorIntentChecker: } for b in self.brand_self: - if _word_boundary_match(normalized, b): + if word_boundary_match(normalized, b): entities["brand_self"].append(b) evidence.append({"type": "entity", "key": "brand_self", "value": b, "match": b}) diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py index 9124a98ac36..d7ce3438fca 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_end_user_permission/mcp_end_user_permission.py @@ -193,7 +193,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail): MCPRequestHandler, ) - access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(mcp_access_groups) + access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups(mcp_access_groups) return list(set(direct_mcp_servers + access_group_servers)) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index a4e50cc183a..a60059fcc4d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -19,7 +19,7 @@ from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast import aiohttp -from pydantic import ConfigDict, TypeAdapter, with_config +from pydantic import ConfigDict, JsonValue, TypeAdapter, with_config from typing_extensions import NotRequired, ReadOnly, TypedDict import litellm @@ -279,7 +279,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self, presidio_analyzer_api_base: str | None = None, presidio_anonymizer_api_base: str | None = None, - ): + ) -> None: self.presidio_analyzer_api_base: str | None = presidio_analyzer_api_base or get_secret( "PRESIDIO_ANALYZER_API_BASE", None ) @@ -922,7 +922,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): def raise_exception_if_blocked_entities_detected( self, analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse - ): + ) -> None: """ Raise an exception if blocked entities are detected """ @@ -1022,7 +1022,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): cache: DualCache, data: dict, call_type: str, - ): + ) -> dict[str, object]: """ - Check if request turned off pii - Check if user allowed to turn off pii (key permissions -> 'allow_pii_controls') @@ -1212,10 +1212,10 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): async def async_post_call_success_hook( self, - data: dict, + data: dict[str, object], user_api_key_dict: UserAPIKeyAuth, response: ModelResponse | EmbeddingResponse | ImageResponse, - ): + ) -> dict[str, JsonValue] | ModelResponse | EmbeddingResponse | ImageResponse: """ Output parse the response object to replace the masked tokens with user sent values """ @@ -1546,7 +1546,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): and delta.get("type") == "text_delta" and isinstance(delta.get("text"), str) ): - unmasked = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(delta["text"], pii_tokens) + unmasked = OPTIONAL_PresidioPIIMasking._unmask_pii_text(delta["text"], pii_tokens) if unmasked != delta["text"]: event["delta"]["text"] = unmasked line = "data: " + json.dumps(event, ensure_ascii=False) @@ -1707,7 +1707,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): return None - def print_verbose(self, print_statement): + def print_verbose(self, print_statement) -> None: try: verbose_proxy_logger.debug(print_statement) if litellm.set_verbose: @@ -1771,3 +1771,6 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self.presidio_analyze_chunk_size_bytes = self._coerce_analyze_chunk_size( litellm_params.presidio_analyze_chunk_size_bytes ) + + +OPTIONAL_PresidioPIIMasking = _OPTIONAL_PresidioPIIMasking diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index bacbd728d89..192d55b0b91 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -134,7 +134,7 @@ def _presidio_output_mode(mode: str | list[str] | Mode, *, include_mcp: bool) -> def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> tuple[CustomGuardrail, ...]: from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) explicit_filter_scope: Final = litellm_params.presidio_filter_scope @@ -163,7 +163,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> params.update(overrides) # Passed outside the heterogeneous params dict so the argument keeps # its precise int | None type. - callback: Final = _OPTIONAL_PresidioPIIMasking( + callback: Final = OPTIONAL_PresidioPIIMasking( presidio_analyze_chunk_size_bytes=litellm_params.presidio_analyze_chunk_size_bytes, **params, ) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 7b2ad86f3dc..d14cc940885 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -36,8 +36,9 @@ from litellm.proxy.guardrails.guardrail_hooks.grayswan import ( ) from litellm.proxy.guardrails.guardrail_hooks.lakera_ai import lakeraAI_Moderation from litellm.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import LakeraAIGuardrail -from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, +from litellm.proxy.guardrails.guardrail_hooks.presidio import ( # noqa: F401 # legacy module exports + OPTIONAL_PresidioPIIMasking, + _OPTIONAL_PresidioPIIMasking, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) from litellm.proxy.guardrails.guardrail_hooks.tool_permission import ( ToolPermissionGuardrail, @@ -226,7 +227,7 @@ guardrail_class_registry: Final[dict[str, type[CustomGuardrail]]] = { SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail, SupportedGuardrailIntegrations.LAKERA.value: lakeraAI_Moderation, SupportedGuardrailIntegrations.LAKERA_V2.value: LakeraAIGuardrail, - SupportedGuardrailIntegrations.PRESIDIO.value: _OPTIONAL_PresidioPIIMasking, + SupportedGuardrailIntegrations.PRESIDIO.value: OPTIONAL_PresidioPIIMasking, SupportedGuardrailIntegrations.TOOL_PERMISSION.value: ToolPermissionGuardrail, } diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 468dc7f19bb..ce30334b8e9 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -166,7 +166,7 @@ def _get_random_llm_message(): return [{"role": "user", "content": random.choice(messages)}] -def _clean_endpoint_data(endpoint_data: dict, details: bool | None = True): +def clean_endpoint_data(endpoint_data: Mapping[str, object], details: bool | None = True) -> dict[str, object]: """ Keep only the explicitly approved, JSON-safe diagnostic fields for display to users. """ @@ -174,6 +174,9 @@ def _clean_endpoint_data(endpoint_data: dict, details: bool | None = True): return {k: v for k, v in endpoint_data.items() if k in displayed} +_clean_endpoint_data: Final = clean_endpoint_data + + def health_check_filter_kwargs_from_general_settings( general_settings: dict | None, ) -> dict: @@ -543,7 +546,9 @@ async def _run_model_health_check(model: dict): model_info, litellm_params, # any-ok: untyped router config dict ) - litellm_params = _update_litellm_params_for_health_check(model_info, litellm_params) + litellm_params = update_litellm_params_for_health_check( # rebind-ok: pre-existing rebinding on a rename-only line + model_info, litellm_params + ) timeout: Final = model_info.get("health_check_timeout") or HEALTH_CHECK_TIMEOUT_SECONDS return await run_with_timeout( @@ -649,12 +654,12 @@ async def _perform_health_check( _model_id = (model.get("model_info") or {}).get("id") if isinstance(is_healthy, dict) and "error" not in is_healthy: - cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details) + cleaned = clean_endpoint_data({**litellm_params, **is_healthy}, details) if _model_id: cleaned["model_id"] = _model_id healthy_endpoints.append(cleaned) elif isinstance(is_healthy, dict): - cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details) + cleaned = clean_endpoint_data({**litellm_params, **is_healthy}, details) if _model_id: cleaned["model_id"] = _model_id if "exception" in is_healthy: @@ -665,7 +670,7 @@ async def _perform_health_check( cleaned["exception_status"] = getattr(exc, "status_code", 500) unhealthy_endpoints.append(cleaned) else: - cleaned = _clean_endpoint_data(litellm_params, details) + cleaned = clean_endpoint_data(litellm_params, details) if _model_id: cleaned["model_id"] = _model_id if isinstance(is_healthy, Exception): @@ -772,7 +777,7 @@ def _resolve_health_check_max_tokens(model_info: dict, litellm_params: dict) -> return None -def _update_litellm_params_for_health_check(model_info: dict, litellm_params: dict) -> dict: +def update_litellm_params_for_health_check(model_info: dict, litellm_params: dict) -> dict: """ Update the litellm params for health check. @@ -865,6 +870,9 @@ def _update_litellm_params_for_health_check(model_info: dict, litellm_params: di return litellm_params +_update_litellm_params_for_health_check: Final = update_litellm_params_for_health_check + + async def perform_health_check( model_list: list, model: str | None = None, diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index ff12f75479a..4180ca1a656 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -53,15 +53,17 @@ from litellm.proxy.db.health_check_latest import ( query_latest_health_checks, ) from litellm.proxy.db.proxy_worker_heartbeat import count_live_proxy_workers -from litellm.proxy.health_check import ( +from litellm.proxy.health_check import ( # noqa: F401 # legacy module exports ADMIN_ONLY_HEALTH_DISPLAY_PARAMS, - _clean_endpoint_data, - _update_litellm_params_for_health_check, + _clean_endpoint_data, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _update_litellm_params_for_health_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + clean_endpoint_data, deployments_targeted_by_name, health_check_filter_kwargs_from_general_settings, perform_health_check, resolve_health_check_mode, run_with_timeout, + update_litellm_params_for_health_check, ) from litellm.proxy.middleware.admission_control_middleware import ( get_admission_control_stats, @@ -634,7 +636,7 @@ async def health_services_endpoint( ) -def _convert_health_check_to_dict(check) -> dict: +def convert_health_check_to_dict(check) -> dict: """Convert health check database record to dictionary format""" return { "health_check_id": check.health_check_id, @@ -652,6 +654,9 @@ def _convert_health_check_to_dict(check) -> dict: } +_convert_health_check_to_dict: Final = convert_health_check_to_dict + + def _check_prisma_client(): """Helper to check if prisma_client is available and raise appropriate error""" from litellm.proxy.proxy_server import prisma_client @@ -883,7 +888,7 @@ async def _save_health_check_results_if_changed( return all(row is not None for row in rows) -async def _save_background_health_checks_to_db( +async def save_background_health_checks_to_db( prisma_client, model_list: list, healthy_endpoints: list, @@ -941,6 +946,9 @@ async def _save_background_health_checks_to_db( return False +_save_background_health_checks_to_db: Final = save_background_health_checks_to_db + + _PROXY_ADMIN_ROLES: Final = frozenset( { LitellmUserRoles.PROXY_ADMIN.value, @@ -1353,7 +1361,7 @@ async def health_check_history_endpoint( ) # Convert to dict format for JSON response using helper function - history_data: Final = [_convert_health_check_to_dict(check) for check in history] + history_data: Final = [convert_health_check_to_dict(check) for check in history] return { "health_checks": history_data, @@ -1385,7 +1393,7 @@ async def latest_health_checks_endpoint( # Convert to dict format for JSON response using helper function checks_data: Final = { - (check.model_id if check.model_id else check.model_name): _convert_health_check_to_dict(check) + (check.model_id if check.model_id else check.model_name): convert_health_check_to_dict(check) for check in latest_checks } @@ -2234,7 +2242,7 @@ async def test_model_connection( stored_params=_OBJECT_MAPPING.validate_python(config_litellm_params), request_params=_OBJECT_MAPPING.validate_python(request_litellm_params), ) - litellm_params = _update_litellm_params_for_health_check( + litellm_params = update_litellm_params_for_health_check( model_info=dict(probe_model_info), litellm_params=litellm_params, ) @@ -2272,7 +2280,7 @@ async def test_model_connection( ) # Clean the result for display - cleaned_result: Final = _clean_endpoint_data({**litellm_params, **result}, details=True) + cleaned_result: Final = clean_endpoint_data({**litellm_params, **result}, details=True) return { "status": "error" if "error" in result else "success", diff --git a/litellm/proxy/hooks/__init__.py b/litellm/proxy/hooks/__init__.py index 0e78a0843cd..138050ccc9a 100644 --- a/litellm/proxy/hooks/__init__.py +++ b/litellm/proxy/hooks/__init__.py @@ -3,35 +3,53 @@ from typing import Final, Literal from . import * from .autorouter_baseline_cache import AutoRouterBaselineCache -from .cache_control_check import _PROXY_CacheControlCheck +from .cache_control_check import ( # noqa: F401 # backwards-compatible package export + PROXY_CacheControlCheck, + _PROXY_CacheControlCheck, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) from .litellm_skills import SkillsInjectionHook -from .max_budget_per_session_limiter import _PROXY_MaxBudgetPerSessionHandler -from .max_iterations_limiter import _PROXY_MaxIterationsHandler -from .parallel_request_limiter import _PROXY_MaxParallelRequestsHandler -from .parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 +from .max_budget_per_session_limiter import ( # noqa: F401 # backwards-compatible package export + PROXY_MaxBudgetPerSessionHandler, + _PROXY_MaxBudgetPerSessionHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) +from .max_iterations_limiter import ( # noqa: F401 # backwards-compatible package export + PROXY_MaxIterationsHandler, + _PROXY_MaxIterationsHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) +from .parallel_request_limiter import ( # noqa: F401 # backwards-compatible package export + PROXY_MaxParallelRequestsHandler, + _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) +from .parallel_request_limiter_v3 import ( # noqa: F401 # backwards-compatible package export + PROXY_MaxParallelRequestsHandler_v3, + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) from .prompt_cache_prediction import PromptCacheObserver from .responses_id_security import ResponsesIDSecurity -from .sensitive_data_routing import _PROXY_SensitiveDataRoutingHandler +from .sensitive_data_routing import ( # noqa: F401 # backwards-compatible package export + PROXY_SensitiveDataRoutingHandler, + _PROXY_SensitiveDataRoutingHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export +) # List of all available hooks that can be enabled. # Defined before the enterprise import below so that any module re-imported # transitively through `enterprise.enterprise_hooks` can resolve `PROXY_HOOKS` # and `get_proxy_hook` from this partially-initialized module without circling. PROXY_HOOKS: Final = { - "parallel_request_limiter": _PROXY_MaxParallelRequestsHandler_v3, - "cache_control_check": _PROXY_CacheControlCheck, + "parallel_request_limiter": PROXY_MaxParallelRequestsHandler_v3, + "cache_control_check": PROXY_CacheControlCheck, "responses_id_security": ResponsesIDSecurity, "litellm_skills": SkillsInjectionHook, - "max_iterations_limiter": _PROXY_MaxIterationsHandler, - "max_budget_per_session_limiter": _PROXY_MaxBudgetPerSessionHandler, - "sensitive_data_routing": _PROXY_SensitiveDataRoutingHandler, + "max_iterations_limiter": PROXY_MaxIterationsHandler, + "max_budget_per_session_limiter": PROXY_MaxBudgetPerSessionHandler, + "sensitive_data_routing": PROXY_SensitiveDataRoutingHandler, "prompt_cache_prediction": PromptCacheObserver, "autorouter_baseline_cache": AutoRouterBaselineCache, } ## FEATURE FLAG HOOKS ## if os.getenv("LEGACY_MULTI_INSTANCE_RATE_LIMITING", "false").lower() == "true": - PROXY_HOOKS["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler + PROXY_HOOKS["parallel_request_limiter"] = PROXY_MaxParallelRequestsHandler def get_proxy_hook( diff --git a/litellm/proxy/hooks/azure_content_safety.py b/litellm/proxy/hooks/azure_content_safety.py index ad3ec844fac..7d0db84692c 100644 --- a/litellm/proxy/hooks/azure_content_safety.py +++ b/litellm/proxy/hooks/azure_content_safety.py @@ -77,7 +77,7 @@ class _PROXY_AzureContentSafety( return result - async def test_violation(self, content: str, source: str | None = None): + async def test_violation(self, content: str, source: str | None = None) -> None: verbose_proxy_logger.debug("Testing Azure Content-Safety for: %s", content) # Construct a request @@ -115,7 +115,7 @@ class _PROXY_AzureContentSafety( cache: DualCache, data: dict, call_type: str, # "completion", "embeddings", "image_generation", "moderation" - ): + ) -> None: verbose_proxy_logger.debug("Inside Azure Content-Safety Pre-Call Hook") try: if is_text_content_call_type(call_type): @@ -135,7 +135,7 @@ class _PROXY_AzureContentSafety( data: dict, user_api_key_dict: UserAPIKeyAuth, response, - ): + ) -> None: verbose_proxy_logger.debug("Inside Azure Content-Safety Post-Call Hook") if not isinstance(response, litellm.ModelResponse): return @@ -148,10 +148,13 @@ class _PROXY_AzureContentSafety( if isinstance(content, str): await self.test_violation(content=content, source="output") - # async def async_post_call_streaming_hook( - # self, - # user_api_key_dict: UserAPIKeyAuth, - # response: str, - # ): - # verbose_proxy_logger.debug("Inside Azure Content-Safety Call-Stream Hook") - # await self.test_violation(content=response, source="output") + +PROXY_AzureContentSafety: Final = _PROXY_AzureContentSafety + +# async def async_post_call_streaming_hook( +# self, +# user_api_key_dict: UserAPIKeyAuth, +# response: str, +# ): +# verbose_proxy_logger.debug("Inside Azure Content-Safety Call-Stream Hook") +# await self.test_violation(content=response, source="output") diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index c4591eb17fe..2a1719ed4c4 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -148,12 +148,12 @@ def resolve_batch_enqueued_token_scopes( def canonical_provider_batch_id(batch_id: str) -> str: from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper get_batch_id_from_unified_batch_id, get_original_file_id, + is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper ) - decoded: Final = _is_base64_encoded_unified_file_id(batch_id) + decoded: Final = is_base64_encoded_unified_file_id(batch_id) if isinstance(decoded, str): if "llm_batch_id" in decoded or "generic_response_id" in decoded: return get_batch_id_from_unified_batch_id(decoded) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 55bd622651e..17f8bd277dd 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -68,15 +68,15 @@ if TYPE_CHECKING: from opentelemetry.trace import Span as _Span from litellm.caching.caching import DualCache + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + PROXY_MaxParallelRequestsHandler_v3 as _ParallelRequestLimiter, + ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( RateLimitDescriptor as _RateLimitDescriptor, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( RateLimitStatus as _RateLimitStatus, ) - from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3 as _ParallelRequestLimiter, - ) from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache from litellm.router import Router as _Router from litellm.types.llms.openai import HttpxBinaryResponseContent @@ -164,16 +164,16 @@ class _PROXY_BatchRateLimiter(CustomLogger): return None from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, decode_model_from_file_id, get_models_from_unified_file_id, + is_base64_encoded_unified_file_id, ) model_from_file_id: Final = decode_model_from_file_id(input_file_id) if model_from_file_id: return model_from_file_id - unified_file_id: Final = _is_base64_encoded_unified_file_id(input_file_id) + unified_file_id: Final = is_base64_encoded_unified_file_id(input_file_id) if unified_file_id: target_model_names: Final = get_models_from_unified_file_id(unified_file_id) if target_model_names: @@ -250,7 +250,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): minute with the submission. The daily descriptor uses its own key so its 24h window never collides with the online limiter's counters. """ - descriptors: Final = self.parallel_request_limiter._create_rate_limit_descriptors( + descriptors: Final = self.parallel_request_limiter.create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data=data, rpm_limit_type=None, @@ -892,13 +892,13 @@ class _PROXY_BatchRateLimiter(CustomLogger): try: # Check if this is a managed file (base64 encoded unified file ID) from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, get_models_from_unified_file_id, + is_base64_encoded_unified_file_id, ) # Managed files require bypassing the HTTP endpoint (which runs access-check hooks) # and calling the managed files hook directly with the user's credentials. - is_managed_file: Final = _is_base64_encoded_unified_file_id(file_id) + is_managed_file: Final = is_base64_encoded_unified_file_id(file_id) # For managed files the unified file id encodes the proxy model # alias(es) the file was uploaded for; auth validates against those. target_model_names: Final = get_models_from_unified_file_id(is_managed_file) if is_managed_file else [] @@ -1039,11 +1039,11 @@ class _PROXY_BatchRateLimiter(CustomLogger): enforces on `/chat/completions` apply here. """ from litellm.proxy.auth.auth_checks import ( - _check_team_member_model_access, - _key_access_group_grants_model, can_key_call_model, can_team_access_model, + check_team_member_model_access, get_team_object, + key_access_group_grants_model, ) from litellm.proxy.proxy_server import llm_router, prisma_client, proxy_logging_obj, user_api_key_cache @@ -1092,14 +1092,14 @@ class _PROXY_BatchRateLimiter(CustomLogger): except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: raise - if not await _key_access_group_grants_model( + if not await key_access_group_grants_model( model=model_to_check, valid_token=user_api_key_dict, team_object=team_object, llm_router=llm_router, ): raise - await _check_team_member_model_access( + await check_team_member_model_access( model=model_to_check, team_object=team_object, valid_token=user_api_key_dict, @@ -1281,3 +1281,6 @@ class _PROXY_BatchRateLimiter(CustomLogger): verbose_proxy_logger.error("Error in batch rate limiting: %s", e, exc_info=True) # Don't block the request if rate limiting fails return data + + +PROXY_BatchRateLimiter = _PROXY_BatchRateLimiter diff --git a/litellm/proxy/hooks/batch_redis_get.py b/litellm/proxy/hooks/batch_redis_get.py index 94ad6fb7cb2..94eb0bb52d1 100644 --- a/litellm/proxy/hooks/batch_redis_get.py +++ b/litellm/proxy/hooks/batch_redis_get.py @@ -25,7 +25,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): self.async_get_cache ) # map the litellm 'get_cache' function to our custom function - def print_verbose(self, print_statement, debug_level: Literal["INFO", "DEBUG"] = "DEBUG"): + def print_verbose(self, print_statement, debug_level: Literal["INFO", "DEBUG"] = "DEBUG") -> None: if debug_level == "DEBUG" or debug_level == "INFO": verbose_proxy_logger.debug(print_statement) if litellm.set_verbose is True: @@ -37,7 +37,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): cache: DualCache, data: dict, call_type: str, - ): + ) -> None: try: """ Get the user key @@ -83,7 +83,7 @@ class _PROXY_BatchRedisRequests(CustomLogger): ) verbose_proxy_logger.debug(traceback.format_exc()) - async def async_get_cache(self, *args, **kwargs): + async def async_get_cache(self, *args, **kwargs) -> object | None: """ - Check if the cache key is in-memory @@ -113,3 +113,6 @@ class _PROXY_BatchRedisRequests(CustomLogger): return litellm.cache.get_cache_logic(cached_result=cached_result, max_age=max_age) except Exception: return None + + +PROXY_BatchRedisRequests: Final = _PROXY_BatchRedisRequests diff --git a/litellm/proxy/hooks/cache_control_check.py b/litellm/proxy/hooks/cache_control_check.py index dab2ed0b933..22af491f2e7 100644 --- a/litellm/proxy/hooks/cache_control_check.py +++ b/litellm/proxy/hooks/cache_control_check.py @@ -24,7 +24,7 @@ class _PROXY_CacheControlCheck(CustomLogger): cache: DualCache, data: dict, call_type: str, - ): + ) -> None: try: verbose_proxy_logger.debug("Inside Cache Control Check Pre-Call Hook") allowed_cache_controls: Final = user_api_key_dict.allowed_cache_controls @@ -56,3 +56,6 @@ class _PROXY_CacheControlCheck(CustomLogger): verbose_logger.exception( "litellm.proxy.hooks.cache_control_check.py::async_pre_call_hook(): Exception occured - %s", e ) + + +PROXY_CacheControlCheck: Final = _PROXY_CacheControlCheck diff --git a/litellm/proxy/hooks/dynamic_rate_limiter.py b/litellm/proxy/hooks/dynamic_rate_limiter.py index 5c99a1bacc1..646afab73d6 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter.py @@ -1,7 +1,6 @@ # What is this? ## Allocates dynamic tpm/rpm quota for a project based on current traffic ## Tracks num active projects per minute - import asyncio import os from collections.abc import Callable @@ -23,7 +22,7 @@ from litellm.proxy.hooks.rate_limiter_utils import ( resolve_llm_provider_for_rate_limit, ) from litellm.types.router import ModelGroupInfo -from litellm.types.utils import CallTypesLiteral +from litellm.types.utils import CallTypesLiteral, LLMResponseTypes from litellm.utils import get_utc_datetime @@ -83,7 +82,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): def __init__(self, internal_usage_cache: DualCache, time_fn: Callable[[], datetime] = get_utc_datetime): self.internal_usage_cache = DynamicRateLimiterCache(cache=internal_usage_cache, time_fn=time_fn) - def update_variables(self, llm_router: Router): + def update_variables(self, llm_router: Router) -> None: self.llm_router = llm_router @with_service_target("rate_limits") @@ -241,7 +240,9 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): return None @with_service_target("rate_limits") - async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes + ) -> LLMResponseTypes | None: try: if isinstance(response, ModelResponse): model_id: Final = response.hidden_params["model_id"] @@ -281,3 +282,6 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): "litellm.proxy.hooks.dynamic_rate_limiter.py::async_post_call_success_hook(): Exception occured - %s", e ) return response + + +PROXY_DynamicRateLimitHandler: Final = _PROXY_DynamicRateLimitHandler diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index 0a078321d5e..d174d9fecb1 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -20,11 +20,12 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import ( ProxyRateLimitError, map_v3_rate_limit_type, ) -from litellm.proxy.hooks.parallel_request_limiter_v3 import ( +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( # noqa: F401 # legacy module exports + PROXY_MaxParallelRequestsHandler_v3, RateLimitDescriptor, RateLimitDescriptorRateLimitObject, RateLimitResponse, - _PROXY_MaxParallelRequestsHandler_v3, + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export claim_request_stash_for_data, get_or_create_request_stash, ) @@ -38,7 +39,7 @@ from litellm.router_utils.add_retry_fallback_headers import ( response_has_hidden_params, ) from litellm.types.router import ModelGroupInfo -from litellm.types.utils import CallTypesLiteral +from litellm.types.utils import CallTypesLiteral, LLMResponseTypes if TYPE_CHECKING: from litellm.types.utils import PriorityReservationSettings @@ -90,9 +91,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): time_provider: Callable[[], datetime] | None = None, ): self.internal_usage_cache = InternalUsageCache(dual_cache=internal_usage_cache) - self.v3_limiter = _PROXY_MaxParallelRequestsHandler_v3(self.internal_usage_cache, time_provider=time_provider) + self.v3_limiter = PROXY_MaxParallelRequestsHandler_v3(self.internal_usage_cache, time_provider=time_provider) - def update_variables(self, llm_router: Router): + def update_variables(self, llm_router: Router) -> None: self.llm_router = llm_router def _get_saturation_check_cache_ttl(self) -> int: @@ -659,7 +660,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): return None @with_service_target("rate_limits") - async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response + ) -> LLMResponseTypes: """ Post-call hook to add rate limit headers to response. Leverages v3 limiter's post-call hook functionality. @@ -689,7 +692,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): return response @with_service_target("rate_limits") - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: """ Update token usage for priority-based rate limiting after successful API calls. @@ -804,3 +807,6 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): except Exception as e: verbose_proxy_logger.exception("Error in dynamic rate limiter success event: %s", e) + + +PROXY_DynamicRateLimitHandlerV3: Final = _PROXY_DynamicRateLimitHandlerV3 diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index a1100474671..ef93809d47b 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -21,7 +21,10 @@ from litellm.proxy._types import ( UpdateKeyRequest, UserAPIKeyAuth, ) -from litellm.proxy.utils import _hash_token_if_needed +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports + _hash_token_if_needed, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + hash_token_if_needed, +) from litellm.secret_managers.base_secret_manager import BaseSecretManager if TYPE_CHECKING: @@ -140,7 +143,7 @@ class KeyManagementEventHooks: ), changed_by_api_key=user_api_key_dict.api_key, table_name=LitellmTableNames.KEY_TABLE_NAME, - object_id=_hash_token_if_needed(data.key), + object_id=hash_token_if_needed(data.key), action="updated", updated_values=json.dumps(updated_fields, default=str), before_value=json.dumps(existing_key_row.json(exclude_none=True), default=str), diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 2ba42c43dcc..cf3025d8c7b 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -130,7 +130,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): return None @with_service_target("session_budgets") - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: """ After a successful LLM call, increment the session spend by the response cost. """ @@ -271,3 +271,6 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): local_only=True, ) return new_value + + +PROXY_MaxBudgetPerSessionHandler: Final = _PROXY_MaxBudgetPerSessionHandler diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index efcafc1b6b0..5b7d7d04e24 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -221,3 +221,6 @@ class _PROXY_MaxIterationsHandler(CustomLogger): local_only=True, ) return new_value + + +PROXY_MaxIterationsHandler: Final = _PROXY_MaxIterationsHandler diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py index 15c28d64d05..d33811d3dea 100644 --- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py +++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py @@ -169,7 +169,7 @@ class SemanticToolFilterHook(CustomLogger): def _selected_tool_names(self, filtered_tools: Sequence[object]) -> list[str]: """Names of the semantically selected tools, as produced by the MCP expansion.""" - names: Final = (self.filter._extract_tool_info(tool)[0] for tool in filtered_tools) + names: Final = (self.filter.extract_tool_info(tool)[0] for tool in filtered_tools) return [name for name in names if name] @staticmethod @@ -219,7 +219,7 @@ class SemanticToolFilterHook(CustomLogger): return False if isinstance(tool, dict) and tool.get("type") == "function" and isinstance(tool.get("name"), str): return False - name, _ = self.filter._extract_tool_info(tool) + name, _ = self.filter.extract_tool_info(tool) return bool(name) and name in self.filter._tool_map def _get_metadata_variable_name(self, data: dict) -> str: @@ -397,14 +397,14 @@ class SemanticToolFilterHook(CustomLogger): filtered_mcp_names: Final[set[str]] = set() for t in filtered_mcp_tools: - name, _ = self.filter._extract_tool_info(t) + name, _ = self.filter.extract_tool_info(t) if name: filtered_mcp_names.add(name) filtered_tools: Final[list[object]] = [] for i, t in enumerate(tools): if i in mcp_indices: - name, _ = self.filter._extract_tool_info(t) + name, _ = self.filter.extract_tool_info(t) if name in filtered_mcp_names: filtered_tools.append(t) else: diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index e810b98f336..8bca51784f2 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -493,7 +493,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): return healthy_deployments @with_service_target("model_budgets") - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: """ Track spend for virtual key + model in DualCache @@ -627,3 +627,6 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): key=marker_key, value=1, ttl=ttl_seconds, refresh_ttl=True ) return polls == 1 + + +PROXY_VirtualKeyModelMaxBudgetLimiter = _PROXY_VirtualKeyModelMaxBudgetLimiter diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index facae4cbb81..0f884387673 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -57,7 +57,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): def __init__(self, internal_usage_cache: InternalUsageCache): self.internal_usage_cache = internal_usage_cache - def print_verbose(self, print_statement): + def print_verbose(self, print_statement) -> None: try: verbose_proxy_logger.debug(print_statement) if litellm.set_verbose: @@ -253,7 +253,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): cache: DualCache, data: dict, call_type: str, - ): + ) -> None: self.print_verbose("Inside Max Parallel Request Pre-Call Hook") api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict) max_parallel_requests = user_api_key_dict.max_parallel_requests @@ -494,7 +494,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) @with_service_target("rate_limits") - async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time) -> None: from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, ) @@ -700,7 +700,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): self.print_verbose(e) @with_service_target("rate_limits") - async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None: 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) @@ -808,7 +808,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): return None @with_service_target("rate_limits") - async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): + async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response) -> None: """ Retrieve the key's remaining rate limits. """ @@ -861,3 +861,6 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) return await super().async_post_call_success_hook(data, user_api_key_dict, response) + + +PROXY_MaxParallelRequestsHandler: Final = _PROXY_MaxParallelRequestsHandler diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 1a3452d9f87..5b865f3b5c3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -867,10 +867,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if self._batch_rate_limiter is None: try: from litellm.proxy.hooks.batch_rate_limiter import ( - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) - self._batch_rate_limiter = _PROXY_BatchRateLimiter( + self._batch_rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=self.internal_usage_cache, parallel_request_limiter=self, time_provider=self._time_provider, @@ -1022,17 +1022,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ if not isinstance(data, dict): return - base_capped_floor: Final = _PROXY_MaxParallelRequestsHandler_v3.no_max_tokens_output_floor(min_configured_limit) + base_capped_floor: Final = PROXY_MaxParallelRequestsHandler_v3.no_max_tokens_output_floor(min_configured_limit) capped_floor: Final = ( max(base_capped_floor, RESPONSES_API_MIN_OUTPUT_TOKENS) if call_type in RESPONSES_API_CALL_TYPES else base_capped_floor ) baseline_floor: Final = DEFAULT_MAX_TOKENS_ESTIMATE // _TPM_FLOOR_FRACTION - is_embedding: Final = _PROXY_MaxParallelRequestsHandler_v3._is_embedding_request(data, call_type) + is_embedding: Final = PROXY_MaxParallelRequestsHandler_v3._is_embedding_request(data, call_type) if ( capped_floor >= baseline_floor - or _PROXY_MaxParallelRequestsHandler_v3._has_explicit_output_cap(data, call_type) + or PROXY_MaxParallelRequestsHandler_v3._has_explicit_output_cap(data, call_type) or is_embedding or endpoint_type == EndpointType.DECISIONS ): @@ -3188,7 +3188,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if (limit := tag_limits.get(tag)) is not None ) - def _create_rate_limit_descriptors( + def create_rate_limit_descriptors( self, user_api_key_dict: UserAPIKeyAuth, data: dict, @@ -3337,6 +3337,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return descriptors + _create_rate_limit_descriptors = create_rate_limit_descriptors + async def _check_model_has_recent_failures( self, model: str, @@ -3943,7 +3945,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if requested_model and self._is_dynamic_rate_limiting_enabled(rpm_limit_type, tpm_limit_type) else False ) - descriptors: Final = self._create_rate_limit_descriptors( # pyright: ignore[reportUnknownMemberType] # legacy helper reads a dictionary with validated keys + descriptors: Final = self.create_rate_limit_descriptors( # pyright: ignore[reportUnknownMemberType] # legacy helper reads a dictionary with validated keys user_api_key_dict=user_api_key_dict, data=dict(data), rpm_limit_type=rpm_limit_type, @@ -4031,7 +4033,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data: dict, call_type: str, endpoint_type: EndpointType = EndpointType.GENERIC, - ): + ) -> Exception | str | dict[str, object] | None: """ Pre-call hook to check rate limits before making the API call. Supports dynamic rate limiting based on deployment health. @@ -5048,7 +5050,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return pipeline_operations @with_service_target("rate_limits") - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: """ Update TPM usage on successful API calls by incrementing counters using pipeline """ @@ -5166,7 +5168,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) @with_service_target("rate_limits") - async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None: """ On failure: decrement max_parallel_requests and refund the upfront TPM reservation only against the scopes the reservation actually @@ -5295,7 +5297,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): await self._release_stashed_parallel_slot(get_request_stash(), None) @with_service_target("rate_limits") - async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): + async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response) -> None: """ Release completed-request slots and update rate limit headers in the response. """ @@ -5476,3 +5478,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): except Exception as e: verbose_proxy_logger.exception("Error releasing TPM reservation on post-call failure: %s", e) return + + +PROXY_MaxParallelRequestsHandler_v3 = _PROXY_MaxParallelRequestsHandler_v3 diff --git a/litellm/proxy/hooks/prompt_injection_detection.py b/litellm/proxy/hooks/prompt_injection_detection.py index 3c2eefcc933..0fee5b061f6 100644 --- a/litellm/proxy/hooks/prompt_injection_detection.py +++ b/litellm/proxy/hooks/prompt_injection_detection.py @@ -76,7 +76,7 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): "and start from scratch", ] - def print_verbose(self, print_statement, level: Literal["INFO", "DEBUG"] = "DEBUG"): + def print_verbose(self, print_statement, level: Literal["INFO", "DEBUG"] = "DEBUG") -> None: if level == "INFO": verbose_proxy_logger.info(print_statement) elif level == "DEBUG": @@ -85,7 +85,7 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): if litellm.set_verbose is True: print(print_statement) # noqa: T201 - def update_environment(self, router: Router | None = None): + def update_environment(self, router: Router | None = None) -> None: self.llm_router = router if self.prompt_injection_params is not None and self.prompt_injection_params.llm_api_check is True: @@ -150,9 +150,9 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, - data: dict, + data: dict[str, object], call_type: str, # "completion", "embeddings", "image_generation", "moderation" - ): + ) -> dict[str, object] | str | None: try: """ - check if user id part of call @@ -278,3 +278,6 @@ class _OPTIONAL_PromptInjectionDetection(CustomLogger): ) return is_prompt_attack + + +OPTIONAL_PromptInjectionDetection = _OPTIONAL_PromptInjectionDetection diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index ffe752329d4..196e0f20ad3 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -45,9 +45,10 @@ from litellm.proxy.spend_tracking.spend_log_error_logger import ( should_suppress_spend_log_tracebacks, spend_log_error, ) -from litellm.proxy.spend_tracking.spend_tracking_utils import ( - _sanitize_error_information_for_spend_logs, +from litellm.proxy.spend_tracking.spend_tracking_utils import ( # noqa: F401 # legacy module exports + _sanitize_error_information_for_spend_logs, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export get_request_model_access_groups, + sanitize_error_information_for_spend_logs, should_store_prompts_and_responses_in_spend_logs, ) from litellm.proxy.utils import ProxyUpdateSpend @@ -152,7 +153,7 @@ class _ProxyDBLogger(CustomLogger): original_exception: Exception, user_api_key_dict: UserAPIKeyAuth, traceback_str: str | None = None, - ): + ) -> None: try: await _release_budget_reservation(budget_reservation=user_api_key_dict.budget_reservation) except Exception: @@ -168,7 +169,7 @@ class _ProxyDBLogger(CustomLogger): request_route: Final = user_api_key_dict.request_route if ( - _ProxyDBLogger._should_track_errors_in_db() is False + ProxyDBLogger._should_track_errors_in_db() is False or request_route is not None and not ( RouteChecks.is_llm_api_route(route=request_route) or RouteChecks.is_info_route(route=request_route) @@ -197,11 +198,11 @@ class _ProxyDBLogger(CustomLogger): # here because the input above is constructed non-None. _error_information = cast( StandardLoggingPayloadErrorInformation, - _sanitize_error_information_for_spend_logs(_error_information, original_exception=original_exception), + sanitize_error_information_for_spend_logs(_error_information, original_exception=original_exception), ) _metadata["error_information"] = _error_information - _metadata = await _ProxyDBLogger._enrich_failure_metadata_unless_db_stalled( + _metadata = await ProxyDBLogger._enrich_failure_metadata_unless_db_stalled( # rebind-ok: pre-existing rebinding on a rename-only line metadata=_metadata, original_exception=original_exception ) @@ -316,7 +317,7 @@ class _ProxyDBLogger(CustomLogger): # Only fetch key details when user_id wasn't already populated (e.g. direct MCP REST calls). # Avoids a cache/DB lookup on every normal LLM request. if metadata.get("user_api_key") and not metadata.get("user_api_key_user_id"): - metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( # rebind-ok: enriched metadata replaces the original + metadata = await ProxyDBLogger._enrich_failure_metadata_with_key_info( # rebind-ok: enriched metadata replaces the original metadata=metadata, resolve_missing_key_identity=str(kwargs.get("call_type")) not in _CAPTURED_IDENTITY_CALL_TYPES, ) @@ -503,7 +504,7 @@ class _ProxyDBLogger(CustomLogger): ) -> dict[str, object]: if isinstance(original_exception, DBLookupDeadlineExceeded): return metadata - return await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) + return await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) @staticmethod async def _enrich_failure_metadata_with_key_info(metadata: dict, resolve_missing_key_identity: bool = True) -> dict: @@ -594,6 +595,9 @@ class _ProxyDBLogger(CustomLogger): return +ProxyDBLogger: Final = _ProxyDBLogger + + def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: patch = {k: v for k, v in metadata.items() if (k.startswith("user_api_key") or k == "tags") and v is not None} if not patch: @@ -609,7 +613,7 @@ def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: async def run_spend_event(line: bytes) -> None: - await _ProxyDBLogger().run_spend_event(line) + await ProxyDBLogger().run_spend_event(line) def _is_unbilled_interaction_response(completion_response: object) -> bool: diff --git a/litellm/proxy/hooks/sensitive_data_routing.py b/litellm/proxy/hooks/sensitive_data_routing.py index 1773fc2d50a..0909f36dbe5 100644 --- a/litellm/proxy/hooks/sensitive_data_routing.py +++ b/litellm/proxy/hooks/sensitive_data_routing.py @@ -204,3 +204,6 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger): data["metadata"] = metadata return data + + +PROXY_SensitiveDataRoutingHandler: Final = _PROXY_SensitiveDataRoutingHandler diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index cd64363f3d3..28fb9baf2aa 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -280,11 +280,11 @@ async def image_edit_api( # The validation will be done at the model level if image is truly required from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -298,7 +298,7 @@ async def image_edit_api( # Read request body and convert UploadFiles to BytesIO ######################################################### form_fields: Final = coerce_numeric_form_fields( - parsed_body=await _read_request_body(request=request), + parsed_body=await read_request_body(request=request), numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS, ) data: Final = { @@ -345,7 +345,7 @@ async def image_edit_api( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 452c653eef0..03a93016177 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -65,7 +65,10 @@ from litellm.proxy.common_utils.callback_utils import ( get_metadata_variable_name_from_kwargs, strip_callback_config, ) -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) from litellm.proxy.spend_tracking.carried_budget_state import carried_budget_metadata from litellm.types.integrations.anthropic_cache_control_hook import GATEWAY_INJECTED_CACHE_METADATA_KEY @@ -639,7 +642,7 @@ def _strip_router_reserved_metadata( ) -def _get_metadata_variable_name(request: Request) -> str: +def get_metadata_variable_name(request: Request) -> str: """ Helper to return what the "metadata" field should be called in the request data @@ -653,6 +656,9 @@ def _get_metadata_variable_name(request: Request) -> str: return metadata_variable_name_for_route(get_request_route(request)) +_get_metadata_variable_name: Final = get_metadata_variable_name + + def metadata_variable_name_for_route(route: str) -> Literal["metadata", "litellm_metadata"]: if "thread" in route or "assistant" in route: return "litellm_metadata" @@ -941,7 +947,7 @@ def convert_key_logging_metadata_to_callback( return team_callback_settings_obj -def _get_validated_callback_metadata(item: dict, *, source: str) -> AddTeamCallback | None: +def get_validated_callback_metadata(item: dict, *, source: str) -> AddTeamCallback | None: try: return AddTeamCallback(**item) except (PydanticValidationError, ValueError) as e: @@ -953,6 +959,9 @@ def _get_validated_callback_metadata(item: dict, *, source: str) -> AddTeamCallb return None +_get_validated_callback_metadata: Final = get_validated_callback_metadata + + class KeyAndTeamLoggingSettings: """ Helper class to get the dynamic logging settings for the key and team @@ -971,7 +980,7 @@ class KeyAndTeamLoggingSettings: return None -def _get_dynamic_logging_metadata( +def get_dynamic_logging_metadata( user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig ) -> TeamCallbackMetadata | None: callback_settings_obj: TeamCallbackMetadata | None = None @@ -986,7 +995,7 @@ def _get_dynamic_logging_metadata( ######################################################################################### if key_dynamic_logging_settings is not None: for item in key_dynamic_logging_settings: - callback = _get_validated_callback_metadata(item=item, source="key-level") + callback = get_validated_callback_metadata(item=item, source="key-level") if callback is None: continue callback_settings_obj = convert_key_logging_metadata_to_callback( @@ -998,7 +1007,7 @@ def _get_dynamic_logging_metadata( ######################################################################################### elif team_dynamic_logging_settings is not None: for item in team_dynamic_logging_settings: - callback = _get_validated_callback_metadata(item=item, source="team-level") + callback = get_validated_callback_metadata(item=item, source="team-level") if callback is None: continue callback_settings_obj = convert_key_logging_metadata_to_callback( @@ -1032,6 +1041,9 @@ def _get_dynamic_logging_metadata( return callback_settings_obj +_get_dynamic_logging_metadata: Final = get_dynamic_logging_metadata + + _TENANT_OTEL_PARAMS: Final = TypeAdapter(StandardCallbackDynamicParams) @@ -1123,7 +1135,7 @@ def resolve_tenant_otel_destinations( callbacks: Final = tuple( callback for item in entries - if (callback := _get_validated_callback_metadata(item=item, source="otel-destination")) is not None + if (callback := get_validated_callback_metadata(item=item, source="otel-destination")) is not None if callback.callback_name.lower() not in disabled ) return tuple( @@ -1440,7 +1452,7 @@ class LiteLLMProxyRequestSetup: """ Add headers to the LLM call by model group """ - from litellm.proxy.auth.auth_checks import _check_model_access_helper + from litellm.proxy.auth.auth_checks import check_model_access_helper from litellm.proxy.proxy_server import llm_router data_model: Final = data.get("model") @@ -1449,7 +1461,7 @@ class LiteLLMProxyRequestSetup: data_model is not None and litellm.model_group_settings is not None and litellm.model_group_settings.forward_client_headers_to_llm_api is not None - and _check_model_access_helper( + and check_model_access_helper( model=data_model, llm_router=llm_router, models=litellm.model_group_settings.forward_client_headers_to_llm_api, @@ -1745,7 +1757,7 @@ class LiteLLMProxyRequestSetup: ## KEY-LEVEL SPEND LOGS / TAGS if "tags" in key_metadata and key_metadata["tags"] is not None: - data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags( + data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=data[_metadata_variable_name].get("tags"), tags_to_add=key_metadata["tags"], ) @@ -1793,8 +1805,8 @@ class LiteLLMProxyRequestSetup: team_spend_logs_metadata=team_metadata.get("spend_logs_metadata"), request_spend_logs_metadata=metadata.get("spend_logs_metadata"), ) - tags: Final = LiteLLMProxyRequestSetup._merge_tags( - request_tags=LiteLLMProxyRequestSetup._merge_tags( + tags: Final = LiteLLMProxyRequestSetup.merge_tags( + request_tags=LiteLLMProxyRequestSetup.merge_tags( request_tags=request_tags if isinstance(request_tags, list) else None, tags_to_add=team_tags if isinstance(team_tags, list) else None, ), @@ -1826,7 +1838,7 @@ class LiteLLMProxyRequestSetup: return {**(team_values or {}), **(request_values or {})} @staticmethod - def _merge_tags(request_tags: list | None, tags_to_add: list | None) -> list: + def merge_tags(request_tags: list | None, tags_to_add: list | None) -> list: """ Helper function to merge two lists of tags, ensuring no duplicates. @@ -1849,6 +1861,8 @@ class LiteLLMProxyRequestSetup: return final_tags + _merge_tags = merge_tags + @staticmethod def add_team_based_callbacks_from_config( team_id: str, @@ -1938,7 +1952,7 @@ class LiteLLMProxyRequestSetup: metadata: Final = _normalized_metadata_slot(request_data, _metadata_variable_name) existing_tags: Final = metadata.get("tags") - metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags( + metadata["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=existing_tags if isinstance(existing_tags, list) else None, tags_to_add=key_tags, ) @@ -1968,7 +1982,7 @@ class LiteLLMProxyRequestSetup: # No allow_client_tags opt-in: caller-supplied tags always flow # into metadata.tags (see add_litellm_data_to_request). The pre-auth # merge mirrors that so _tag_max_budget_check sees the same tags. - headers: Final = _safe_get_request_headers(request=request) + headers: Final = safe_get_request_headers(request=request) raw_header_tags: Final = headers.get("x-litellm-tags") if not raw_header_tags: return @@ -1990,7 +2004,7 @@ class LiteLLMProxyRequestSetup: metadata: Final = _normalized_metadata_slot(request_data, _metadata_variable_name) existing_tags: Final = metadata.get("tags") - metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags( + metadata["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=existing_tags if isinstance(existing_tags, list) else None, tags_to_add=header_tags, ) @@ -2088,7 +2102,7 @@ async def add_litellm_data_to_request( if _mk.startswith("user_api_key_"): del _user_metadata[_mk] - _raw_headers: Final[dict[str, str]] = RedactedDict(_safe_get_request_headers(request)) + _raw_headers: Final[dict[str, str]] = RedactedDict(safe_get_request_headers(request)) forward_llm_auth = False if general_settings: @@ -2162,7 +2176,7 @@ async def add_litellm_data_to_request( } safe_add_api_version_from_query_params(data, request) - _metadata_variable_name: Final = _get_metadata_variable_name(request) + _metadata_variable_name: Final = get_metadata_variable_name(request) if data.get(_metadata_variable_name, None) is None: data[_metadata_variable_name] = {} @@ -2494,7 +2508,7 @@ async def add_litellm_data_to_request( ) if tags is not None: - data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags( + data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=data[_metadata_variable_name].get("tags"), tags_to_add=tags, ) @@ -2506,7 +2520,7 @@ async def add_litellm_data_to_request( else None ) if _caller_body_tags: - data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags( # rebind-ok: matches file idiom + data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup.merge_tags( # rebind-ok: matches file idiom request_tags=data[_metadata_variable_name].get("tags"), tags_to_add=_caller_body_tags, ) @@ -2523,7 +2537,7 @@ async def add_litellm_data_to_request( ) # Team Callbacks controls - callback_settings_obj: Final = _get_dynamic_logging_metadata( + callback_settings_obj: Final = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) if callback_settings_obj is not None: @@ -3004,7 +3018,7 @@ def _enforced_params_check( return True -def _add_guardrails_from_key_or_team_metadata( +def add_guardrails_from_key_or_team_metadata( key_metadata: dict | None, team_metadata: dict | None, data: dict, @@ -3024,7 +3038,7 @@ def _add_guardrails_from_key_or_team_metadata( project_metadata: The project metadata dictionary to check for guardrails """ - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check # Initialize guardrails set (avoiding duplicates) combined_guardrails: Final = set() @@ -3032,19 +3046,19 @@ def _add_guardrails_from_key_or_team_metadata( # Add key-level guardrails first if key_metadata and "guardrails" in key_metadata: if isinstance(key_metadata["guardrails"], list) and len(key_metadata["guardrails"]) > 0: - _premium_user_check() + premium_user_check() combined_guardrails.update(key_metadata["guardrails"]) # Add team-level guardrails (set automatically handles duplicates) if team_metadata and "guardrails" in team_metadata: if isinstance(team_metadata["guardrails"], list) and len(team_metadata["guardrails"]) > 0: - _premium_user_check() + premium_user_check() combined_guardrails.update(team_metadata["guardrails"]) # Add project-level guardrails (set automatically handles duplicates) if project_metadata and "guardrails" in project_metadata: if isinstance(project_metadata["guardrails"], list) and len(project_metadata["guardrails"]) > 0: - _premium_user_check() + premium_user_check() combined_guardrails.update(project_metadata["guardrails"]) # Set combined guardrails in metadata as list @@ -3052,6 +3066,9 @@ def _add_guardrails_from_key_or_team_metadata( data[metadata_variable_name]["guardrails"] = list(combined_guardrails) +_add_guardrails_from_key_or_team_metadata: Final = add_guardrails_from_key_or_team_metadata + + def _add_guardrails_from_policies_in_metadata( key_metadata: dict | None, team_metadata: dict | None, @@ -3077,7 +3094,7 @@ def _add_guardrails_from_policies_in_metadata( from litellm._logging import verbose_proxy_logger from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.proxy.policy_engine.policy_resolver import PolicyResolver - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check from litellm.types.proxy.policy_engine import PolicyMatchContext # Collect policy names from key and team metadata @@ -3086,19 +3103,19 @@ def _add_guardrails_from_policies_in_metadata( # Add key-level policies first if key_metadata and "policies" in key_metadata: if isinstance(key_metadata["policies"], list) and len(key_metadata["policies"]) > 0: - _premium_user_check() + premium_user_check() policy_names.update(key_metadata["policies"]) # Add team-level policies if team_metadata and "policies" in team_metadata: if isinstance(team_metadata["policies"], list) and len(team_metadata["policies"]) > 0: - _premium_user_check() + premium_user_check() policy_names.update(team_metadata["policies"]) # Add project-level policies if project_metadata and "policies" in project_metadata: if isinstance(project_metadata["policies"], list) and len(project_metadata["policies"]) > 0: - _premium_user_check() + premium_user_check() policy_names.update(project_metadata["policies"]) if not policy_names: @@ -3166,7 +3183,7 @@ def add_guardrails_from_auth_metadata( metadata_variable_name: str, ) -> None: """Resolve key, team, and project guardrails, direct and via policies, onto the request metadata.""" - _add_guardrails_from_key_or_team_metadata( + add_guardrails_from_key_or_team_metadata( key_metadata=user_api_key_dict.metadata, team_metadata=user_api_key_dict.team_metadata, project_metadata=user_api_key_dict.project_metadata, diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index 97311a0ef8a..3f05c3a407a 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -17,11 +17,15 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry -from litellm.proxy.auth.auth_checks import ( - _cache_access_object, - _cache_key_object, - _cache_team_object, - _get_team_object_from_cache, +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _cache_access_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _cache_team_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_team_object_from_cache, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + cache_access_object, + cache_key_object, + cache_team_object, + get_team_object_from_cache, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET @@ -166,9 +170,9 @@ def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None: def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None: """Admin Viewer parity: PROXY_ADMIN or PROXY_ADMIN_VIEW_ONLY may read.""" - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={"error": CommonProxyErrors.not_allowed_access.value}, @@ -303,7 +307,7 @@ async def _cache_access_group_record(record: _AccessGroupRecord) -> None: from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache access_group_table: Final = _record_to_access_group_table(record) - await _cache_access_object( + await cache_access_object( access_group_id=record.access_group_id, access_group_table=access_group_table, user_api_key_cache=user_api_key_cache, @@ -408,7 +412,7 @@ async def _patch_team_caches_add_access_group( ) -> None: """Patch cached team objects to include access_group_id.""" for team_id in team_ids: - cached_team = await _get_team_object_from_cache( + cached_team = await get_team_object_from_cache( key=f"team_id:{team_id}", user_api_key_cache=user_api_key_cache, parent_otel_span=None, @@ -421,7 +425,7 @@ async def _patch_team_caches_add_access_group( cached_team.access_group_ids = list(cached_team.access_group_ids) + [access_group_id] else: continue - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=cached_team, user_api_key_cache=user_api_key_cache, @@ -437,14 +441,14 @@ async def _patch_team_caches_remove_access_group( ) -> None: """Patch cached team objects to remove access_group_id.""" for team_id in team_ids: - cached_team = await _get_team_object_from_cache( + cached_team = await get_team_object_from_cache( key=f"team_id:{team_id}", user_api_key_cache=user_api_key_cache, parent_otel_span=None, ) if cached_team is not None and cached_team.access_group_ids: cached_team.access_group_ids = [ag for ag in cached_team.access_group_ids if ag != access_group_id] - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=cached_team, user_api_key_cache=user_api_key_cache, @@ -473,7 +477,7 @@ async def _patch_key_caches_add_access_group( cached_key.access_group_ids = list(cached_key.access_group_ids) + [access_group_id] else: continue - await _cache_key_object( + await cache_key_object( hashed_token=token, user_api_key_obj=cached_key, user_api_key_cache=user_api_key_cache, @@ -496,7 +500,7 @@ async def _patch_key_caches_remove_access_group( ) if cached_key is not None and cached_key.access_group_ids: cached_key.access_group_ids = [ag for ag in cached_key.access_group_ids if ag != access_group_id] - await _cache_key_object( + await cache_key_object( hashed_token=token, user_api_key_obj=cached_key, user_api_key_cache=user_api_key_cache, diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index de2a534a21d..ec1996a9839 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -28,9 +28,10 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_checks import ( - _virtual_key_max_budget_check, +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _virtual_key_max_budget_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export can_key_call_resolved_model, + virtual_key_max_budget_check, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.db.autorouter_session_rollup import ( @@ -340,7 +341,7 @@ async def _authorize_models_this_test_can_call( ) try: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, ) @@ -558,10 +559,10 @@ async def preview_auto_router_routing( if member_team is not None and _models_this_test_can_call(resolved.complexity_router_config): from litellm.proxy.auth.user_api_key_auth import ( - _run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse the serving admission policy + run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse the serving admission policy ) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=actor, request=http_request, request_data=request_data, diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index e16ea4a812e..d001dc666af 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -22,8 +22,9 @@ from fastapi import APIRouter, Depends, HTTPException from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time -from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, validate_budget_duration, ) from litellm.proxy.utils import jsonify_object @@ -256,7 +257,7 @@ async def budget_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, @@ -317,7 +318,7 @@ async def list_budget( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, diff --git a/litellm/proxy/management_endpoints/cache_settings_endpoints.py b/litellm/proxy/management_endpoints/cache_settings_endpoints.py index 2bd4b6b47f9..aaf5d30d2e1 100644 --- a/litellm/proxy/management_endpoints/cache_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/cache_settings_endpoints.py @@ -467,7 +467,7 @@ async def get_cache_settings( if prisma_client is not None: cache_config = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"}) if cache_config is not None and cache_config.cache_settings: - stored = proxy_config._decrypt_db_variables( + stored = proxy_config.decrypt_db_variables( # rebind-ok: pre-existing rebinding on a rename-only line variables_dict=_parse_stored_settings(cache_config.cache_settings) ) @@ -534,8 +534,10 @@ async def test_cache_connection( try: existing_row: Final = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"}) if existing_row is not None and existing_row.cache_settings: - saved_settings = proxy_config._decrypt_db_variables( - variables_dict=_parse_stored_settings(existing_row.cache_settings) + saved_settings = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables( + variables_dict=_parse_stored_settings(existing_row.cache_settings) + ) ) except Exception: # noqa: BLE001 - a saved-settings lookup failure must not block a connection test saved_settings = {} @@ -614,7 +616,9 @@ async def update_cache_settings( saved_settings: dict[str, object] = {} if existing_row is not None and existing_row.cache_settings: before_settings = _parse_stored_settings(existing_row.cache_settings) - saved_settings = proxy_config._decrypt_db_variables(variables_dict=before_settings) + saved_settings = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables(variables_dict=before_settings) + ) action: Final[AUDIT_ACTIONS] = "updated" if existing_row is not None else "created" # Preserve stored secrets behind any redacted or omitted credential, then @@ -622,7 +626,7 @@ async def update_cache_settings( cache_settings: Final = _resolve_cache_url_precedence(_merge_over_saved(request.cache_settings, saved_settings)) # Encrypt sensitive fields (keep redis_type for storage) - encrypted_settings: Final = proxy_config._encrypt_env_variables(environment_variables=cache_settings) + encrypted_settings: Final = proxy_config.encrypt_env_variables(environment_variables=cache_settings) # Save to database await _cache_config_table(prisma_client).upsert( @@ -640,13 +644,13 @@ async def update_cache_settings( # Reinitialize cache with new settings # Decrypt for initialization - decrypted_settings: Final = proxy_config._decrypt_db_variables(variables_dict=encrypted_settings) + decrypted_settings: Final = proxy_config.decrypt_db_variables(variables_dict=encrypted_settings) # Remove redis_type if present (UI-only field, not a Cache parameter) cache_params: Final = {k: v for k, v in decrypted_settings.items() if k != "redis_type"} # Initialize cache (frontend sends type="redis", not redis_type) - proxy_config._init_cache(cache_params=cache_params) + proxy_config.init_cache(cache_params=cache_params) # Update the last cache params to avoid reinitializing unnecessarily CacheSettingsManager.update_cache_params(cache_params) diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 6d458feef57..a9385c637e8 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -43,7 +43,7 @@ def validate_budget_duration(budget_duration: str | None, status_code: int = 400 from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache -from litellm.proxy._types import ( +from litellm.proxy._types import ( # re-exported CommonProxyErrors, KeyRequestBase, LiteLLM_ManagementEndpoint_MetadataFields, @@ -56,13 +56,14 @@ from litellm.proxy._types import ( NewProjectRequest, UpdateProjectRequest, UserAPIKeyAuth, -) -from litellm.proxy._types import ( # noqa: F401 re-exported - user_api_key_has_admin_view as _user_has_admin_view, + user_api_key_has_admin_view, ) from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.management.teams.authz import is_team_admin -from litellm.proxy.utils import _premium_user_check +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports + _premium_user_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + premium_user_check, +) from litellm.repositories.team_repository import TeamRepository from litellm.types.utils import BudgetConfig @@ -70,6 +71,8 @@ if TYPE_CHECKING: from litellm.proxy._types import NewProjectRequest, UpdateProjectRequest from litellm.proxy.utils import PrismaClient, ProxyLogging +_user_has_admin_view: Final = user_api_key_has_admin_view + # TODO: drop once the litellm-enterprise pin moves past 0.1.71, which imports this name _is_user_team_admin: Final = is_team_admin @@ -160,7 +163,7 @@ def _passthrough_routes_permission_error(field: str, entity: str) -> HTTPExcepti ) -def _check_passthrough_routes_caller_permission( +def check_passthrough_routes_caller_permission( data: BaseModel | None, user_api_key_dict: UserAPIKeyAuth, *, @@ -178,6 +181,9 @@ def _check_passthrough_routes_caller_permission( ) +_check_passthrough_routes_caller_permission: Final = check_passthrough_routes_caller_permission + + def check_allowed_passthrough_routes_caller_permission( data: BaseModel | None, user_api_key_dict: UserAPIKeyAuth, @@ -235,7 +241,7 @@ def _metadata_changes_denied_routes(data: BaseModel, metadata: object, existing_ return metadata is None and "metadata" in data.model_fields_set and existing_denied is not None -def _check_disable_global_guardrails_caller_permission( +def check_disable_global_guardrails_caller_permission( disable_global_guardrails: bool | None, metadata: Mapping[str, object] | None, user_api_key_dict: UserAPIKeyAuth, @@ -263,7 +269,10 @@ def _check_disable_global_guardrails_caller_permission( ) -def _team_member_has_permission( +_check_disable_global_guardrails_caller_permission: Final = check_disable_global_guardrails_caller_permission + + +def team_member_has_permission( user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable, permission: str, @@ -279,7 +288,10 @@ def _team_member_has_permission( return False -async def _user_has_admin_privileges( +_team_member_has_permission: Final = team_member_has_permission + + +async def user_has_admin_privileges( user_api_key_dict: UserAPIKeyAuth, prisma_client: Optional["PrismaClient"] = None, user_api_key_cache: Optional["DualCache"] = None, @@ -345,6 +357,9 @@ async def _user_has_admin_privileges( return False +_user_has_admin_privileges: Final = user_has_admin_privileges + + def _org_admin_can_invite_user( admin_user_obj: LiteLLM_UserTable, target_user_obj: LiteLLM_UserTable, @@ -487,7 +502,7 @@ async def admin_can_invite_user( return False -def _set_object_metadata_field( +def set_object_metadata_field( object_data: Union[ LiteLLM_TeamTable, KeyRequestBase, @@ -508,12 +523,15 @@ def _set_object_metadata_field( value: Value to set for the field """ if field_name in LiteLLM_ManagementEndpoint_MetadataFields_Premium and value: - _premium_user_check(field_name) + premium_user_check(field_name) object_data.metadata = object_data.metadata or {} object_data.metadata[field_name] = value +_set_object_metadata_field: Final = set_object_metadata_field + + _TEAM_MEMBER_BUDGET_LIMIT_FIELDS: Final = ( "max_budget", "soft_budget", @@ -573,7 +591,7 @@ def _has_meaningful_budget_limit(budget_values: Mapping[str, object]) -> bool: return any(_is_set_budget_value(budget_values.get(field)) for field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS) -async def _upsert_budget_and_membership( +async def upsert_budget_and_membership( tx, *, team_id: str, @@ -583,7 +601,7 @@ async def _upsert_budget_and_membership( budget_patch: Mapping[str, object], team_default_budget_id: str | None = None, shared_budget_ids: frozenset[str] | None = None, -): +) -> None: """ Apply a merge-patch of per-member budget fields to a team membership. @@ -694,6 +712,9 @@ async def _upsert_budget_and_membership( ) +_upsert_budget_and_membership: Final = upsert_budget_and_membership + + def _update_metadata_field(updated_kv: dict, field_name: str) -> None: """ Helper function to update metadata fields that require premium user checks in the update endpoint @@ -708,7 +729,7 @@ def _update_metadata_field(updated_kv: dict, field_name: str) -> None: # only for a truthy value. The falsy value is still persisted below so a # previously-set field can be cleared. if updated_kv.get(field_name): - _premium_user_check() + premium_user_check() if field_name in updated_kv and updated_kv[field_name] is not None: # remove field from updated_kv @@ -730,7 +751,7 @@ def _has_non_empty_value(value: object) -> bool: return True -def _update_metadata_fields(updated_kv: dict) -> None: +def update_metadata_fields(updated_kv: dict) -> None: """ Helper function to update all metadata fields (both premium and standard). @@ -744,3 +765,6 @@ def _update_metadata_fields(updated_kv: dict) -> None: for field in LiteLLM_ManagementEndpoint_MetadataFields: if field in updated_kv and updated_kv[field] is not None: _update_metadata_field(updated_kv=updated_kv, field_name=field) + + +_update_metadata_fields: Final = update_metadata_fields diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index 6b4f2b690db..a9c3f2b953f 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -191,7 +191,7 @@ def _mask_sensitive_fields(data: Mapping[str, object], sensitive_fields: set[str return masked -def _get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, str | None]: +def get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, str | None]: """Read current env var values as fallback when no DB record exists.""" values: Final = {} for field_name, env_var_name in env_var_mapping.items(): @@ -200,6 +200,9 @@ def _get_current_env_values(env_var_mapping: dict[str, str]) -> dict[str, str | return values +_get_current_env_values: Final = get_current_env_values + + class _JsonSchemaField(TypedDict, total=False): type: ReadOnly[str] anyOf: ReadOnly[Sequence["_JsonSchemaField"]] @@ -232,14 +235,17 @@ def _build_field_schema(model_class: type[BaseModel]) -> dict[str, object]: } -def _parse_config_value(raw: str | Mapping[str, object]) -> dict[str, object]: +def parse_config_value(raw: str | Mapping[str, object]) -> dict[str, object]: """Parse a config_value from DB (may be JSON string or dict).""" if isinstance(raw, str): return safe_json_loads(raw, default={}) return dict(raw) -def _set_env_vars( +_parse_config_value: Final = parse_config_value + + +def set_env_vars( config_data: Mapping[str, object], env_var_mapping: Mapping[str, str] = HASHICORP_ENV_VAR_MAPPING, ) -> None: @@ -252,24 +258,30 @@ def _set_env_vars( os.environ.pop(env_var_name, None) -def _clear_hashicorp_vault_state(proxy_config: "ProxyConfig") -> None: +_set_env_vars: Final = set_env_vars + + +def clear_hashicorp_vault_state(proxy_config: "ProxyConfig") -> None: """Clear all Hashicorp Vault state: env vars, secret manager, and change-detection cache.""" - _set_env_vars({}) + set_env_vars({}) if litellm._key_management_system == KeyManagementSystem.HASHICORP_VAULT: litellm.secret_manager_client = None litellm._key_management_system = None proxy_config._last_hashicorp_vault_config = None # pyright: ignore[reportPrivateUsage] # proxy-internal change-detection cache +_clear_hashicorp_vault_state: Final = clear_hashicorp_vault_state + + def _snapshot_cyberark_boot_env(proxy_config: "ProxyConfig") -> None: """Capture deployment-provided CYBERARK_* env vars once, before the first DB-driven overwrite.""" if proxy_config._cyberark_boot_env is None: # pyright: ignore[reportPrivateUsage] # proxy-internal boot snapshot - proxy_config._cyberark_boot_env = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) # pyright: ignore[reportPrivateUsage] # proxy-internal boot snapshot + proxy_config._cyberark_boot_env = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) # pyright: ignore[reportPrivateUsage] # proxy-internal boot snapshot def _restore_cyberark_runtime(proxy_config: "ProxyConfig", env_values: Mapping[str, str | None]) -> None: """Restore CYBERARK_* env vars and reinitialize (or drop) the secret manager to match them.""" - _set_env_vars(env_values, CYBERARK_ENV_VAR_MAPPING) + set_env_vars(env_values, CYBERARK_ENV_VAR_MAPPING) if env_values.get("cyberark_api_base"): try: proxy_config.initialize_secret_manager(key_management_system="cyberark") @@ -367,15 +379,19 @@ async def update_hashicorp_vault_config( existing_decrypted: dict[str, object] | None = None env_values: dict[str, str | None] = {} if existing_record is not None and existing_record.config_value is not None: - existing_data: Final = _parse_config_value(existing_record.config_value) - existing_decrypted = proxy_config._decrypt_db_variables(existing_data) + existing_data: Final = parse_config_value(existing_record.config_value) + existing_decrypted = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables(existing_data) + ) for field in HASHICORP_ENV_VAR_MAPPING: if field not in config_data and existing_decrypted.get(field): config_data[field] = existing_decrypted[field] else: # No DB record (or DB record with null config_value) — merge from # current env vars instead. - env_values = _get_current_env_values(HASHICORP_ENV_VAR_MAPPING) + env_values = get_current_env_values( # rebind-ok: pre-existing rebinding on a rename-only line + HASHICORP_ENV_VAR_MAPPING + ) for field in HASHICORP_ENV_VAR_MAPPING: if field not in config_data and env_values.get(field): config_data[field] = env_values[field] @@ -404,15 +420,15 @@ async def update_hashicorp_vault_config( ) # Snapshot current env vars so we can restore on failure - previous_env: Final = _get_current_env_values(HASHICORP_ENV_VAR_MAPPING) + previous_env: Final = get_current_env_values(HASHICORP_ENV_VAR_MAPPING) # Set env vars and verify the secret manager can initialize before persisting - _set_env_vars(config_data) + set_env_vars(config_data) try: proxy_config.initialize_secret_manager(key_management_system="hashicorp_vault") except Exception as e: - _set_env_vars(previous_env) + set_env_vars(previous_env) verbose_proxy_logger.exception("Error reinitializing Hashicorp Vault secret manager: %s", str(e)) raise HTTPException( status_code=500, @@ -420,7 +436,7 @@ async def update_hashicorp_vault_config( ) # Only persist to DB after successful init - encrypted_data: Final = proxy_config._encrypt_env_variables(config_data) + encrypted_data: Final = proxy_config.encrypt_env_variables(config_data) config_value: Final = safe_dumps(encrypted_data) await _config_overrides_table(prisma_client).upsert( where={"config_type": "hashicorp_vault"}, @@ -474,11 +490,11 @@ async def get_hashicorp_vault_config( Get current Hashicorp Vault configuration. Returns decrypted values from DB, or falls back to current env vars. """ - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view from litellm.proxy.proxy_server import prisma_client, proxy_config # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail="Only admin users can view config overrides", @@ -498,10 +514,10 @@ async def get_hashicorp_vault_config( ) if db_record is not None and db_record.config_value is not None: - config_data: Final = _parse_config_value(db_record.config_value) + config_data: Final = parse_config_value(db_record.config_value) # Decrypt then mask sensitive fields so plaintext secrets are never sent to the UI - decrypted_data: Final[Mapping[str, object]] = proxy_config._decrypt_db_variables(config_data) + decrypted_data: Final[Mapping[str, object]] = proxy_config.decrypt_db_variables(config_data) masked_data: Final = _mask_sensitive_fields(decrypted_data, HASHICORP_SENSITIVE_FIELDS) return ConfigOverrideSettingsResponse( @@ -511,7 +527,7 @@ async def get_hashicorp_vault_config( ) # Fallback to env vars — also mask sensitive values - env_values: Final = _get_current_env_values(HASHICORP_ENV_VAR_MAPPING) + env_values: Final = get_current_env_values(HASHICORP_ENV_VAR_MAPPING) masked_env_values: Final = _mask_sensitive_fields(env_values, HASHICORP_SENSITIVE_FIELDS) return ConfigOverrideSettingsResponse( @@ -556,7 +572,9 @@ async def delete_hashicorp_vault_config( before_config: dict[str, object] | None = None if existing_record is not None and existing_record.config_value is not None: try: - before_config = proxy_config._decrypt_db_variables(_parse_config_value(existing_record.config_value)) + before_config = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables(parse_config_value(existing_record.config_value)) + ) except Exception: before_config = None @@ -568,7 +586,7 @@ async def delete_hashicorp_vault_config( except RecordNotFoundError: verbose_proxy_logger.debug("No existing Hashicorp Vault config record to delete") - _clear_hashicorp_vault_state(proxy_config) + clear_hashicorp_vault_state(proxy_config) # Only emit audit log if a row was actually removed; an idempotent # delete on a non-existent row produces no security-relevant change. @@ -688,13 +706,13 @@ async def update_cyberark_config( existing_decrypted: dict[str, object] | None = None # mutable-ok: DB payload # rebind-ok: set when record exists env_values: dict[str, str | None] = {} # mutable-ok: env snapshot # rebind-ok: populated when no DB record exists if existing_record is not None and existing_record.config_value is not None: - existing_data: Final = _parse_config_value(existing_record.config_value) + existing_data: Final = parse_config_value(existing_record.config_value) existing_decrypted = proxy_config._decrypt_db_variables(existing_data) # pyright: ignore[reportPrivateUsage] # rebind-ok: populated when a prior record decrypts for field in CYBERARK_ENV_VAR_MAPPING: if field not in config_data and existing_decrypted.get(field): config_data[field] = existing_decrypted[field] else: - env_values = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) # rebind-ok: populated when no DB record exists + env_values = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) # rebind-ok: populated when no DB record exists for field in CYBERARK_ENV_VAR_MAPPING: if field not in config_data and env_values.get(field): config_data[field] = env_values[field] @@ -719,13 +737,13 @@ async def update_cyberark_config( ) _snapshot_cyberark_boot_env(proxy_config) - previous_env: Final = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) - _set_env_vars(config_data, CYBERARK_ENV_VAR_MAPPING) + previous_env: Final = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) + set_env_vars(config_data, CYBERARK_ENV_VAR_MAPPING) try: proxy_config.initialize_secret_manager(key_management_system="cyberark") except Exception as e: # noqa: BLE001 # any init failure must roll back env vars - _set_env_vars(previous_env, CYBERARK_ENV_VAR_MAPPING) + set_env_vars(previous_env, CYBERARK_ENV_VAR_MAPPING) verbose_proxy_logger.exception("Error reinitializing CyberArk secret manager: %s", str(e)) raise HTTPException( status_code=500, @@ -776,11 +794,11 @@ async def get_cyberark_config( Sensitive fields are masked before leaving the server. """ from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage + user_api_key_has_admin_view, # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage ) from litellm.proxy.proxy_server import prisma_client, proxy_config - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail="Only admin users can view config overrides", @@ -797,7 +815,7 @@ async def get_cyberark_config( db_record: Final = await _config_overrides_table(prisma_client).find_unique(where={"config_type": "cyberark"}) if db_record is not None and db_record.config_value is not None: - config_data: Final = _parse_config_value(db_record.config_value) + config_data: Final = parse_config_value(db_record.config_value) decrypted_data: Final[Mapping[str, object]] = proxy_config._decrypt_db_variables(config_data) # pyright: ignore[reportPrivateUsage] # proxy-internal helper, mirrors hashicorp endpoint usage masked_data: Final = _mask_sensitive_fields(decrypted_data, CYBERARK_SENSITIVE_FIELDS) @@ -807,7 +825,7 @@ async def get_cyberark_config( field_schema=field_schema, ) - env_values: Final = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) + env_values: Final = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) masked_env_values: Final = _mask_sensitive_fields(env_values, CYBERARK_SENSITIVE_FIELDS) return ConfigOverrideSettingsResponse( @@ -848,7 +866,7 @@ async def delete_cyberark_config( before_config: dict[str, object] | None = None # mutable-ok: audit snapshot # rebind-ok: set when decrypts if existing_record is not None and existing_record.config_value is not None: try: - before_config = proxy_config._decrypt_db_variables(_parse_config_value(existing_record.config_value)) # pyright: ignore[reportPrivateUsage] # rebind-ok: populated when the prior record decrypts + before_config = proxy_config._decrypt_db_variables(parse_config_value(existing_record.config_value)) # pyright: ignore[reportPrivateUsage] # rebind-ok: populated when the prior record decrypts except Exception: # noqa: BLE001 # undecryptable prior config must not block deletion before_config = None # rebind-ok: reset when decryption fails diff --git a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py index fb8b2544d7d..663e22c4312 100644 --- a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py @@ -214,14 +214,14 @@ def _coordination_redis_source(settings: Mapping[str, object] | None) -> Coordin an explicit block wins, else a plain-Redis response-cache backend is borrowed, else the REDIS_* environment fallback applies. """ - from litellm.proxy.proxy_server import _environment_has_redis_connection_target + from litellm.proxy.proxy_server import environment_has_redis_connection_target if settings: return "coordination_redis" cache_backend: Final = litellm.cache.cache if litellm.cache is not None else None if isinstance(cache_backend, (RedisCache, RedisClusterCache)): return "cache_backend" - if _environment_has_redis_connection_target(): + if environment_has_redis_connection_target(): return "environment" return None @@ -416,7 +416,7 @@ async def check_coordination_redis_connection( Builds a throwaway client (never touching global state) and pings it. """ - from litellm.proxy.proxy_server import _build_redis_usage_cache + from litellm.proxy.proxy_server import build_redis_usage_cache _enforce_proxy_admin(user_api_key_dict) @@ -426,7 +426,9 @@ async def check_coordination_redis_connection( redis_cache: RedisCache | None = None try: - redis_cache = _build_redis_usage_cache(params.model_dump(exclude_none=True)) + redis_cache = build_redis_usage_cache( # rebind-ok: pre-existing rebinding on a rename-only line + params.model_dump(exclude_none=True) + ) await asyncio.wait_for(redis_cache.ping(), timeout=_PING_TIMEOUT_SECONDS) return CoordinationRedisTestResponse(status="healthy") except asyncio.TimeoutError: diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index f9c09128f66..3caaf5ba1af 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -40,15 +40,19 @@ from litellm.proxy.db.db_span import db_span if TYPE_CHECKING: from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import PrismaClient -from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _ALGO_AES_GCM, - _ENCRYPTION_ALGORITHM_SETTING, - _V2_GCM_PREFIX, +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports + _ALGO_AES_GCM, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _ENCRYPTION_ALGORITHM_SETTING, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + ALGO_AES_GCM, + ENCRYPTION_ALGORITHM_SETTING, + V2_GCM_PREFIX, SecretMapDecodeError, - _get_salt_key, + _get_salt_key, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export decode_secret_map, decrypt_value_helper, encrypt_value_helper, + get_salt_key, ) ValueClass = Literal["migrated", "legacy", "plaintext", "undecryptable", "not-a-string"] @@ -126,7 +130,7 @@ class MigrationReport: def is_migrated(value: object) -> bool: """True if ``value`` is already an AES-256-GCM (``v2:gcm:``) ciphertext.""" - return isinstance(value, str) and value.startswith(_V2_GCM_PREFIX) + return isinstance(value, str) and value.startswith(V2_GCM_PREFIX) def classify_value(value: object, key: str = "scan") -> ValueClass: @@ -145,7 +149,7 @@ def classify_value(value: object, key: str = "scan") -> ValueClass: return "not-a-string" if value == "": return "plaintext" - if value.startswith(_V2_GCM_PREFIX): + if value.startswith(V2_GCM_PREFIX): return "migrated" decrypted: Final = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) if decrypted is None: @@ -164,7 +168,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object: """ if not isinstance(value, str) or value == "": return value - if value.startswith(_V2_GCM_PREFIX): + if value.startswith(V2_GCM_PREFIX): return value # idempotent: already migrated decrypted: Final = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) if decrypted is None: @@ -197,11 +201,11 @@ def _assert_aes_gate_enabled() -> None: """ from litellm.proxy.proxy_server import general_settings - algo: Final = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING) - if not (isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM): + algo: Final = general_settings.get(ENCRYPTION_ALGORITHM_SETTING) + if not (isinstance(algo, str) and algo.lower() == ALGO_AES_GCM): raise RuntimeError( "Encryption migration requires general_settings.encryption_algorithm: " - f"'{_ALGO_AES_GCM}'. Current value: {algo!r}. Set it before migrating " + f"'{ALGO_AES_GCM}'. Current value: {algo!r}. Set it before migrating " "so re-encrypted values are written in the AES-256-GCM format." ) @@ -436,13 +440,13 @@ def _classify_callback_value(value: object) -> ValueClass: even when run with the AES write gate off. """ from litellm.proxy.common_utils.callback_utils import ( - _CALLBACK_VAR_ENCRYPTED_PREFIX, + CALLBACK_VAR_ENCRYPTED_PREFIX, ) if not isinstance(value, str): return "not-a-string" inner = value - inner = inner.removeprefix(_CALLBACK_VAR_ENCRYPTED_PREFIX) + inner = inner.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX) # rebind-ok: pre-existing rebinding on a rename-only line return classify_value(inner, key="callback") @@ -595,17 +599,17 @@ async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: obje already-v2 / scanned figures. Returns one report per covered location. """ from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) pre: Final = {r.location: r for r in await _scan_covered_tables(prisma_client)} - current_key: Final = _get_salt_key() + current_key: Final = get_salt_key() if current_key is None: raise RuntimeError( "Cannot migrate covered tables: no salt key / master key is set. Set LITELLM_SALT_KEY before migrating." ) - await _rotate_master_key( + await rotate_master_key( prisma_client=cast("PrismaClient", prisma_client), user_api_key_dict=cast("UserAPIKeyAuth", user_api_key_dict), current_master_key=current_key, diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index 650e743b027..d604c58d1f5 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -37,9 +37,10 @@ from litellm.proxy.common_utils.user_api_key_cache import ( ) from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity from litellm.proxy.management_endpoints.common_utils import validate_budget_duration -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export handle_update_object_permission_common, + set_object_permission, ) from litellm.proxy.utils import handle_exception_on_proxy from litellm.repositories.budget_repository import BudgetRepository @@ -469,7 +470,7 @@ async def new_end_user( ## Handle Object Permission - MCP Servers, Vector Stores etc. new_end_user_obj = _STR_OBJECT_DICT.validate_python( - await _set_object_permission( + await set_object_permission( data_json=new_end_user_obj, prisma_client=prisma_client, ) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 1953370be39..2374f7d7aae 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -62,20 +62,23 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( get_daily_activity, raise_public, ) -from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export require_caller_user_id_for_non_admin, + user_api_key_has_admin_view, validate_budget_duration, validate_finite_spend, ) -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _check_permissions_caller_permission, +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + check_permissions_caller_permission, generate_key_helper_fn, prepare_metadata_fields, ) -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export handle_update_object_permission_common, + set_object_permission, ) from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.proxy.utils import handle_exception_on_proxy, hash_password @@ -592,7 +595,7 @@ async def new_user( if data.auto_create_key and isinstance(user_api_key_dict, UserAPIKeyAuth): enforce_batch_limits_are_admin_only(data, None, user_api_key_dict, "key") - _check_permissions_caller_permission( + check_permissions_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) @@ -602,7 +605,9 @@ async def new_user( # Persist the requested grants as their own row and link it, mirroring key/team creation. # generate_key_helper_fn only forwards object_permission_id, so without this the entitlement # the caller sent would be dropped on the floor. - data_json = await _set_object_permission(data_json=data_json, prisma_client=prisma_client) + data_json = await set_object_permission( # rebind-ok: pre-existing rebinding on a rename-only line + data_json=data_json, prisma_client=prisma_client + ) data_json.pop("password", None) teams = data.teams if teams is None: @@ -788,7 +793,7 @@ def _enforce_user_info_access(user_id: str | None, user_api_key_dict: UserAPIKey # Admin-view roles (PROXY_ADMIN and PROXY_ADMIN_VIEW_ONLY) bypass # ownership, mirroring the `/user/info` carve-out that # `RouteChecks.non_proxy_admin_allowed_routes_check` applies upstream. - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return if user_id == user_api_key_dict.user_id: return @@ -969,7 +974,7 @@ async def user_info( raise Exception( "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" ) - if user_id is None and _user_has_admin_view(user_api_key_dict): + if user_id is None and user_api_key_has_admin_view(user_api_key_dict): return await _get_user_info_for_proxy_admin(user_api_key_dict=user_api_key_dict) elif user_id is None: user_id = user_api_key_dict.user_id @@ -1045,7 +1050,7 @@ async def _check_user_info_v2_access( ) # Rule 1: Proxy admins — fetch and return the target row directly - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return await _fetch_target_user() # Rule 2: Self-lookup @@ -1423,9 +1428,9 @@ async def _invalidate_user_spend_counter_if_changed( and not safely subscriptable). """ if non_default_values.get("spend") is not None: - from litellm.proxy.proxy_server import _invalidate_spend_counter + from litellm.proxy.proxy_server import invalidate_spend_counter - await _invalidate_spend_counter(counter_key=f"spend:user:{non_default_values['user_id']}") + await invalidate_spend_counter(counter_key=f"spend:user:{non_default_values['user_id']}") def _clears_object_permission(user_request: UpdateUserRequest) -> bool: @@ -1488,7 +1493,7 @@ async def _update_single_user_helper( if not user_request.user_id and not user_request.user_email: raise ValueError("Either user_id or user_email must be provided") - _check_permissions_caller_permission( + check_permissions_caller_permission( data=user_request, user_api_key_dict=user_api_key_dict, ) @@ -2156,7 +2161,7 @@ async def _authorize_user_list_request( - Org admins: returns comma-separated org IDs scoped to their allowed orgs. - Others: raises 403. """ - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return organization_ids if user_api_key_dict.user_id is None: @@ -2430,7 +2435,7 @@ async def delete_user( - user_ids: List[str] - The list of user id's to be deleted. """ from litellm.proxy.management_endpoints.team_endpoints import ( - _cleanup_members_with_roles, + cleanup_members_with_roles, ) from litellm.proxy.management_helpers.audit_logs import ( get_audit_log_changed_by, @@ -2546,7 +2551,7 @@ async def delete_user( ).table.find_many(where={"team_id": {"in": user_row.teams}}) teams_to_update: list[tuple[str, str]] = [] for team in fetch_all_teams: - removed_team_members, new_team_members = _cleanup_members_with_roles( + removed_team_members, new_team_members = cleanup_members_with_roles( existing_team_row=LiteLLM_TeamTable.model_validate(team.model_dump()), data=TeamMemberDeleteRequest( team_id=team.team_id, @@ -2680,7 +2685,7 @@ async def _resolve_org_filter_for_user_search( if not ui_settings.get("scope_user_search_to_org", False): return None # flag OFF — no filtering - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return None # proxy admin — see everything # Try to resolve org admin memberships @@ -2875,7 +2880,7 @@ async def ui_view_users( def resolve_user_daily_activity_entity_ids( *, user_id: str | None, user_api_key_dict: UserAPIKeyAuth ) -> tuple[str, ...] | None | ScopeDenied: - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return (user_id,) if user_id is not None else None caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict) diff --git a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py index 292cec1346d..531fbdef2bc 100644 --- a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py +++ b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py @@ -17,7 +17,10 @@ from litellm.proxy._types import ( from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, +) from litellm.repositories.table_repositories import JWTKeyMappingRepository router: Final = APIRouter() @@ -310,7 +313,7 @@ async def list_jwt_key_mappings( from litellm.proxy.proxy_server import prisma_client # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException(status_code=403, detail="Only proxy admins can list JWT key mappings") if prisma_client is None: @@ -348,7 +351,7 @@ async def info_jwt_key_mapping( from litellm.proxy.proxy_server import prisma_client # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException(status_code=403, detail="Only proxy admins can get JWT key mapping info") if prisma_client is None: diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 6063ee21250..ac9dfeb136d 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -53,9 +53,10 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_s ) from litellm.proxy._types import * from litellm.proxy._types import Litellm_EntityType, LiteLLM_VerificationToken, hash_token -from litellm.proxy.auth.auth_checks import ( - _delete_cache_key_object, +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _delete_cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export can_team_access_model, + delete_cache_key_object, get_jwt_key_mapping_cache_keys_for_token, get_key_end_user_budget_id, get_org_object, @@ -88,19 +89,25 @@ from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHoo from litellm.proxy.hooks.model_max_budget_limiter import build_model_max_budget_usage from litellm.proxy.management.teams.authz import TEAM_ADMIN_ONLY, TEAM_OR_ORG_ADMIN, is_team_admin from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, - _check_passthrough_routes_caller_permission, - _set_object_metadata_field, - _team_member_has_permission, - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _check_disable_global_guardrails_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _check_passthrough_routes_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _set_object_metadata_field, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _team_member_has_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export check_allowed_passthrough_routes_caller_permission, check_denied_passthrough_routes_caller_permission, + check_disable_global_guardrails_caller_permission, + check_passthrough_routes_caller_permission, + set_object_metadata_field, + team_member_has_permission, + user_api_key_has_admin_view, validate_budget_duration, validate_finite_spend, ) -from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, +from litellm.proxy.management_endpoints.model_management_endpoints import ( # noqa: F401 # legacy module exports + _add_model_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + add_model_to_db, ) from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights from litellm.proxy.management_endpoints.team_admin_field_permissions import ( @@ -114,11 +121,12 @@ from litellm.proxy.management_helpers.access_group_key_sync import ( sync_key_update_access_group_membership, ) from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export attach_object_permission_to_dict, handle_update_object_permission_common, invalidate_cached_object_permissions, + set_object_permission, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, validate_key_vector_stores_against_team, @@ -130,12 +138,16 @@ from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.proxy.search_endpoints.search_tool_registry import rotate_search_tools_master_key from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start from litellm.proxy.spend_tracking.spend_counter_batch import SPEND_COUNTERS_TARGET -from litellm.proxy.spend_tracking.spend_tracking_utils import _is_master_key -from litellm.proxy.utils import ( +from litellm.proxy.spend_tracking.spend_tracking_utils import ( # noqa: F401 # legacy module exports + _is_master_key, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_master_key, +) +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports PrismaClient, ProxyLogging, - _hash_token_if_needed, + _hash_token_if_needed, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export handle_exception_on_proxy, + hash_token_if_needed, is_valid_api_key, ) from litellm.repositories.base_repository import BaseRepository @@ -526,7 +538,7 @@ def _get_user_in_team(team_table: LiteLLM_TeamTableCachedObj, user_id: str | Non return None -def _get_caller_team_role( +def get_caller_team_role( team_table: LiteLLM_TeamTableCachedObj, user_api_key_dict: UserAPIKeyAuth, ) -> Literal["admin", "user"] | None: @@ -536,7 +548,10 @@ def _get_caller_team_role( return None if member is None else member.role -def _calculate_key_rotation_time(rotation_interval: str) -> datetime: +_get_caller_team_role: Final = get_caller_team_role + + +def calculate_key_rotation_time(rotation_interval: str) -> datetime: """ Helper function to calculate the next rotation time for a key based on the rotation interval. @@ -551,6 +566,9 @@ def _calculate_key_rotation_time(rotation_interval: str) -> datetime: return now + timedelta(seconds=interval_seconds) +_calculate_key_rotation_time: Final = calculate_key_rotation_time + + def _set_key_rotation_fields( data: dict, auto_rotate: bool, @@ -583,7 +601,7 @@ def _set_key_rotation_fields( { "auto_rotate": auto_rotate, "rotation_interval": rotation_interval, - "key_rotation_at": _calculate_key_rotation_time(rotation_interval), + "key_rotation_at": calculate_key_rotation_time(rotation_interval), } ) @@ -630,7 +648,7 @@ def _team_key_operation_team_member_check( detail=f"User={assigned_user_id} not assigned to team={team_table.team_id}", ) - caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) + caller_team_role: Final = get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) is_admin: Final = ( user_api_key_dict.user_role is not None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value @@ -1007,7 +1025,7 @@ def _enforce_allowed_routes_update_permission( ) -def _check_permissions_caller_permission( +def check_permissions_caller_permission( data: GenerateRequestBase, user_api_key_dict: UserAPIKeyAuth, ) -> None: @@ -1029,6 +1047,9 @@ def _check_permissions_caller_permission( ) +_check_permissions_caller_permission: Final = check_permissions_caller_permission + + def _check_budget_limits_delegation_ceiling( budget_limits: list[BudgetLimitEntry] | None, delegation_ceiling: float | None, @@ -1324,11 +1345,11 @@ async def _common_key_generation_helper( is_ui_session_team_key=is_ui_session_team_key, team_table=team_table, ) - _check_permissions_caller_permission( + check_permissions_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, _requested_metadata, user_api_key_dict, @@ -1379,7 +1400,7 @@ async def _common_key_generation_helper( # Set Management Endpoint Metadata Fields for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=data, field_name=field, value=getattr(data, field), @@ -1388,7 +1409,7 @@ async def _common_key_generation_helper( for field in LiteLLM_ManagementEndpoint_MetadataFields: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=data, field_name=field, value=getattr(data, field), @@ -1480,7 +1501,7 @@ async def _common_key_generation_helper( for _op_field, _op_default_value in _default_object_permission.items(): _caller_object_permission.setdefault(_op_field, _op_default_value) - data_json = await _set_object_permission( + data_json = await set_object_permission( # rebind-ok: pre-existing rebinding on a rename-only line data_json=data_json, prisma_client=prisma_client, ) @@ -1744,7 +1765,7 @@ async def _check_team_key_limits( # Exclude the key being updated to avoid double-counting its limits. # data.key may be a raw key (sk-...) or a pre-hashed token_id. if isinstance(data, UpdateKeyRequest) and data.key is not None: - hashed_key: Final = _hash_token_if_needed(data.key) + hashed_key: Final = hash_token_if_needed(data.key) keys = [key for key in keys if key.token != hashed_key] check_team_key_model_specific_limits( keys=keys, @@ -1932,7 +1953,7 @@ async def _check_org_key_limits( # Exclude the key being updated to avoid double-counting its limits. # data.key may be a raw key (sk-...) or a pre-hashed token_id. if isinstance(data, UpdateKeyRequest) and data.key is not None: - hashed_key: Final = _hash_token_if_needed(data.key) + hashed_key: Final = hash_token_if_needed(data.key) keys = [key for key in keys if key.token != hashed_key] check_org_key_model_specific_limits( keys=keys, @@ -2088,7 +2109,7 @@ async def generate_key_fn( user_api_key_dict=user_api_key_dict, allowed_routes_was_provided="allowed_routes" in data.model_fields_set, ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) @@ -2262,7 +2283,7 @@ async def generate_service_account_key_fn( user_api_key_dict=user_api_key_dict, allowed_routes_was_provided="allowed_routes" in data.model_fields_set, ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) @@ -2368,10 +2389,10 @@ def prepare_metadata_fields(data: BaseModel, non_default_values: dict, existing_ else: casted_metadata[k] = v if k in LiteLLM_ManagementEndpoint_MetadataFields_Premium: - from litellm.proxy.utils import _premium_user_check + from litellm.proxy.utils import premium_user_check if v: - _premium_user_check(k) + premium_user_check(k) casted_metadata[k] = v except Exception as e: @@ -2441,7 +2462,7 @@ async def _update_key_row_with_soft_budget( existing_key_row: LiteLLM_VerificationToken, changed_by: str, ) -> _KeyUpdateResult: - hashed_token: Final = _hash_token_if_needed(key) + hashed_token: Final = hash_token_if_needed(key) key_where: Final[_KeyRowWhere] = {"token": hashed_token} tx: _KeyUpdateTx async with prisma_client.tx() as tx: @@ -2508,7 +2529,7 @@ async def prepare_key_update_data( # Set Management Endpoint Metadata Fields for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=data, field_name=field, value=getattr(data, field), @@ -2645,7 +2666,7 @@ async def _get_and_validate_existing_key( ) if token is not None: - hashed_token: Final = _hash_token_if_needed(token=token) + hashed_token: Final = hash_token_if_needed(token=token) existing_key_row: Final[LiteLLM_VerificationToken | None] = await _prisma_table( VerificationTokenRepository(prisma_client) @@ -2742,7 +2763,7 @@ async def _process_single_key_update( # Validate max_budget _validate_max_budget(update_key_request.max_budget) - _check_permissions_caller_permission( + check_permissions_caller_permission( data=update_key_request, user_api_key_dict=user_api_key_dict, ) @@ -2754,7 +2775,7 @@ async def _process_single_key_update( prisma_client=prisma_client, ) - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( update_key_request.disable_global_guardrails, update_key_request.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -2885,8 +2906,8 @@ async def _process_single_key_update( ), user_api_key_cache=user_api_key_cache, ) - await _delete_cache_key_object( - hashed_token=_hash_token_if_needed(key_request.key), + await delete_cache_key_object( + hashed_token=hash_token_if_needed(key_request.key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -2895,7 +2916,7 @@ async def _process_single_key_update( # authenticating against the access groups it just lost. await sync_key_update_access_group_membership( prisma_client=prisma_client, - key_token=_hash_token_if_needed(_resolve_token_to_update(data=key_request, existing_key_row=existing_key_row)), + key_token=hash_token_if_needed(_resolve_token_to_update(data=key_request, existing_key_row=existing_key_row)), data=key_request, existing_key_row=existing_key_row, ) @@ -3100,11 +3121,11 @@ async def _validate_update_key_data( user_api_key_dict=user_api_key_dict, ) check_allowed_passthrough_routes_caller_permission(data, user_api_key_dict) - _check_permissions_caller_permission( + check_permissions_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -3570,8 +3591,8 @@ async def update_key_fn( ), user_api_key_cache=user_api_key_cache, ) - await _delete_cache_key_object( - hashed_token=_hash_token_if_needed(key), + await delete_cache_key_object( + hashed_token=hash_token_if_needed(key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -3580,7 +3601,7 @@ async def update_key_fn( # authenticating against the access groups it just lost. await sync_key_update_access_group_membership( prisma_client=prisma_client, - key_token=_hash_token_if_needed(key), + key_token=hash_token_if_needed(key), data=data, existing_key_row=existing_key_row, ) @@ -3588,7 +3609,7 @@ async def update_key_fn( if data.spend is not None: from litellm.proxy.proxy_server import spend_counter_cache - counter_key: Final = f"spend:key:{_hash_token_if_needed(key)}" + counter_key: Final = f"spend:key:{hash_token_if_needed(key)}" spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=data.spend, ttl=60) if spend_counter_cache.redis_cache is not None: try: @@ -3924,7 +3945,7 @@ async def bulk_update_team_keys( hashed_key_ids: Final = [] seen_hashes: Final = set() for k in data.key_ids: - h = _hash_token_if_needed(k) + h = hash_token_if_needed(k) if h in seen_hashes: continue seen_hashes.add(h) @@ -3970,7 +3991,7 @@ async def bulk_update_team_keys( failed_updates: Final[list[FailedKeyUpdate]] = [] for token in requested_tokens: - db_token = _hash_token_if_needed(token) + db_token = hash_token_if_needed(token) try: if db_token not in existing_by_token: raise HTTPException( @@ -4440,7 +4461,7 @@ async def info_key_fn( key = key or user_api_key_dict.api_key hashed_key: str | None = key if key is not None: - hashed_key = _hash_token_if_needed(token=key) + hashed_key = hash_token_if_needed(token=key) # rebind-ok: pre-existing rebinding on a rename-only line live_key_info: Final = await _prisma_table(VerificationTokenRepository(prisma_client)).find_unique( where={"token": hashed_key}, include={"litellm_budget_table": True}, @@ -5012,7 +5033,7 @@ async def delete_verification_tokens( failed_tokens: list = [] try: if prisma_client: - hashed_tokens: Final[list[str]] = [_hash_token_if_needed(token=key) for key in tokens] + hashed_tokens: Final[list[str]] = [hash_token_if_needed(token=key) for key in tokens] tokens = hashed_tokens _keys_being_deleted: Final[list[LiteLLM_VerificationToken]] = cast( # cast-ok: find_many returns a list "list[LiteLLM_VerificationToken]", @@ -5044,7 +5065,7 @@ async def delete_verification_tokens( status_code=status.HTTP_403_FORBIDDEN, detail={"error": "You are not authorized to delete this key"}, ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=authorized_keys, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -5188,7 +5209,7 @@ async def _save_deleted_verification_token_records( await _deleted_verification_token_table(prisma_client).create_many(data=records) -async def _persist_deleted_verification_tokens( +async def persist_deleted_verification_tokens( keys: Sequence[LiteLLM_VerificationToken], prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, @@ -5208,6 +5229,9 @@ async def _persist_deleted_verification_tokens( ) +_persist_deleted_verification_tokens: Final = persist_deleted_verification_tokens + + async def delete_key_aliases( key_aliases: list[str], user_api_key_cache: UserApiKeyCache, @@ -5228,7 +5252,7 @@ async def delete_key_aliases( ) -async def _rotate_master_key( +async def rotate_master_key( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, current_master_key: str, @@ -5266,7 +5290,7 @@ async def _rotate_master_key( reencrypted for model in decrypted_models if ( - reencrypted := await _add_model_to_db( + reencrypted := await add_model_to_db( model_params=Deployment(**model), user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -5299,10 +5323,10 @@ async def _rotate_master_key( environment_variables_dict = _env_vars_param_value(c) if environment_variables_dict: - decrypted_env_vars: Final = proxy_config._decrypt_and_set_db_env_variables( + decrypted_env_vars: Final = proxy_config.decrypt_and_set_db_env_variables( environment_variables=dict[str, str](environment_variables_dict) ) - encrypted_env_vars: Final = proxy_config._encrypt_env_variables( + encrypted_env_vars: Final = proxy_config.encrypt_env_variables( environment_variables=decrypted_env_vars, new_encryption_key=new_master_key, ) @@ -5399,6 +5423,9 @@ async def _rotate_master_key( verbose_proxy_logger.debug("Successfully re-encrypted %s credentials with new master key", len(credentials)) +_rotate_master_key: Final = rotate_master_key + + def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None: from litellm.proxy._types import CommonProxyErrors @@ -5669,7 +5696,7 @@ async def _execute_virtual_key_regeneration( prisma_client=prisma_client, ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=[key_in_db], prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -5702,8 +5729,8 @@ async def _execute_virtual_key_regeneration( user_api_key_cache=user_api_key_cache, ) if hashed_api_key or key: - await _delete_cache_key_object( - hashed_token=_hash_token_if_needed(key), + await delete_cache_key_object( + hashed_token=hash_token_if_needed(key), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -5739,7 +5766,7 @@ def _check_regenerate_guardrail_opt_out( ) -> None: if data is None: return - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -5836,7 +5863,7 @@ async def regenerate_key_fn( allowed_routes_was_provided="allowed_routes" in data.model_fields_set, ) check_allowed_passthrough_routes_caller_permission(data, user_api_key_dict) - _check_permissions_caller_permission( + check_permissions_caller_permission( data=data, user_api_key_dict=user_api_key_dict, ) @@ -5864,7 +5891,7 @@ async def regenerate_key_fn( is_master_key_regeneration: Final = ( data is not None and data.new_master_key is not None - and _is_master_key(api_key=regenerate_target_key, _master_key=master_key) + and is_master_key(api_key=regenerate_target_key, _master_key=master_key) ) if ( @@ -5892,7 +5919,7 @@ async def regenerate_key_fn( detail={"error": "DB not connected. prisma_client is None"}, ) - _is_master_key_valid: Final = _is_master_key(api_key=key, _master_key=master_key) + _is_master_key_valid: Final = is_master_key(api_key=key, _master_key=master_key) if master_key is not None and data and _is_master_key_valid: if data.new_master_key is None: @@ -5900,7 +5927,7 @@ async def regenerate_key_fn( status_code=status.HTTP_400_BAD_REQUEST, detail={"error": "New master key is required."}, ) - await _rotate_master_key( + await rotate_master_key( prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, current_master_key=master_key, @@ -6292,7 +6319,7 @@ async def reset_key_spend_fn( # a later write would re-fetch and re-cache the pre-write row, pinning # that pod to the stale budget_limits/spend for the rest of its own # cache TTL even though the DB is already correct. - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_api_key, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -6324,7 +6351,7 @@ async def validate_key_list_check( key_hash: str | None, prisma_client: PrismaClient, ) -> LiteLLM_UserTable | None: - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return None if user_api_key_dict.user_id is None: @@ -6451,7 +6478,7 @@ def _get_team_ids_with_key_list_permission_from_objects( team.team_id for team in team_objects if not is_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team) - and _team_member_has_permission( + and team_member_has_permission( user_api_key_dict=user_api_key_dict, team_obj=team, permission=KeyManagementRoutes.KEY_LIST.value, @@ -6672,7 +6699,7 @@ async def list_keys( if not user_id and not is_proxy_admin: user_id = user_api_key_dict.user_id - response: Final = await _list_key_helper( + response: Final = await list_key_helper( prisma_client=prisma_client, page=page, size=size, @@ -7074,7 +7101,7 @@ def _build_key_filter_conditions( return combined_where -async def _list_key_helper( +async def list_key_helper( prisma_client: PrismaClient, page: int, size: int, @@ -7259,6 +7286,9 @@ async def _list_key_helper( ) +_list_key_helper: Final = list_key_helper + + def _get_condition_to_filter_out_ui_session_tokens() -> Mapping[str, object]: """ Condition to filter out UI session tokens @@ -7427,7 +7457,7 @@ async def block_key( ) ## UPDATE KEY CACHE - invalidate so next read re-fetches from DB - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -7541,7 +7571,7 @@ async def unblock_key( ) ## UPDATE KEY CACHE - invalidate so next read re-fetches from DB - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/management_endpoints/management_v1/spend_logs.py b/litellm/proxy/management_endpoints/management_v1/spend_logs.py index b122840004c..b93c05c555f 100644 --- a/litellm/proxy/management_endpoints/management_v1/spend_logs.py +++ b/litellm/proxy/management_endpoints/management_v1/spend_logs.py @@ -48,9 +48,9 @@ async def _spend_log_scope_clause( applies, so a dropdown can never offer a value from a row the caller could not open. """ - from litellm.proxy.spend_tracking.spend_management_endpoints import _is_admin_view_safe, read_scope_sql + from litellm.proxy.spend_tracking.spend_management_endpoints import is_admin_view_safe, read_scope_sql - if _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + if is_admin_view_safe(user_api_key_dict=user_api_key_dict): return None, () scope: Final = await resolve_owned_read_scope( user_api_key_dict.user_id, partial(log_team_lookup, user_api_key_dict) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 82022639095..5f65bc498ad 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -176,13 +176,14 @@ if MCP_AVAILABLE: store_user_oauth_credential, update_mcp_server, ) - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _raise_if_not_oauth2, + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: F401 # legacy module exports + _raise_if_not_oauth2, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export authorize_with_server, client_supplied_application_type, client_supplied_redirect_uris, exchange_token_with_server, get_request_base_url, + raise_if_not_oauth2, redeem_passthrough_authorization_code, register_client_with_server, resolve_ephemeral_dcr_client, @@ -225,15 +226,20 @@ if MCP_AVAILABLE: UserMCPManagementMode, is_per_server_oauth_discovery_eligible, ) - from litellm.proxy.auth.user_api_key_auth import ( - _user_api_key_auth_builder, + from litellm.proxy.auth.user_api_key_auth import ( # noqa: F401 # legacy module exports + _user_api_key_auth_builder, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export user_api_key_auth, + user_api_key_auth_builder, ) - from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, + from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export populate_request_with_path_params, + read_request_body, + ) + from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, ) - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.management_endpoints.mcp_connector_import import ( ConnectorConversionError, ConvertedConnector, @@ -913,7 +919,7 @@ if MCP_AVAILABLE: or (payload.auth_type is not None and payload.auth_type != existing.auth_type) ) - def _inherit_credentials_from_existing_server( + def inherit_credentials_from_existing_server( payload: NewMCPServerRequest, ) -> NewMCPServerRequest: if not payload.server_id: @@ -982,6 +988,8 @@ if MCP_AVAILABLE: payload_dict["credentials"] = inherited_credentials return NewMCPServerRequest.model_validate(payload_dict) + _inherit_credentials_from_existing_server: Final = inherit_credentials_from_existing_server + async def _resolve_session_server_id(payload: NewMCPServerRequest) -> str: """Decide the id an OAuth session runs under. @@ -1204,8 +1212,8 @@ if MCP_AVAILABLE: """ from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.management_helpers.object_permission_utils import ( - _get_allow_all_keys_server_ids, - _get_team_allowed_mcp_servers, + get_allow_all_keys_server_ids, + get_team_allowed_mcp_servers, ) from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -1216,8 +1224,8 @@ if MCP_AVAILABLE: check_db_only=True, ) - team_server_ids: Final = await _get_team_allowed_mcp_servers(team_obj) - allow_all_server_ids: Final = _get_allow_all_keys_server_ids() + team_server_ids: Final = await get_team_allowed_mcp_servers(team_obj) + allow_all_server_ids: Final = get_allow_all_keys_server_ids() all_allowed_ids: Final = team_server_ids | allow_all_server_ids if not all_allowed_ids: @@ -1228,7 +1236,7 @@ if MCP_AVAILABLE: for server_id in all_allowed_ids: server = global_mcp_server_manager.get_mcp_server_by_id(server_id) if server is not None: - mcp_server_table = global_mcp_server_manager._build_mcp_server_table(server) + mcp_server_table = global_mcp_server_manager.build_mcp_server_table(server) servers.append(mcp_server_table) return _redact_mcp_credentials_list(servers) @@ -1314,7 +1322,7 @@ if MCP_AVAILABLE: # Only proxy admins may query another team's MCP servers. # Non-admins must belong to the requested team. sanitized_team_id: Final = team_id.strip() - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) if not is_admin: from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.proxy_server import ( @@ -1820,7 +1828,7 @@ if MCP_AVAILABLE: from litellm.proxy.auth.ip_address_utils import IPAddressUtils client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) - is_admin_view: Final = _user_has_admin_view(user_api_key_dict) + is_admin_view: Final = user_api_key_has_admin_view(user_api_key_dict) is_restricted_virtual_key: Final = _is_restricted_virtual_key_request(user_api_key_dict) resolved: Final = await resolve_mcp_server( server_id, @@ -2124,7 +2132,7 @@ if MCP_AVAILABLE: ) created_by: Final = user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME - payload_with_credentials: Final = _inherit_credentials_from_existing_server(payload) + payload_with_credentials: Final = inherit_credentials_from_existing_server(payload) temp_record: Final = _build_temporary_mcp_server_record( payload_with_credentials, created_by, @@ -2229,7 +2237,7 @@ if MCP_AVAILABLE: # grants must NOT bypass auth (see comment above). path_lower: Final = get_request_route(request).rstrip("/").lower() if path_lower.endswith("/token"): - body_data: Final = await _read_request_body(request=request) + body_data: Final = await read_request_body(request=request) grant_type: Final = (body_data or {}).get("grant_type", "") if grant_type != "authorization_code": # Fall through to normal LiteLLM auth (will 401 if @@ -2244,10 +2252,12 @@ if MCP_AVAILABLE: # token can be minted via the redirect alone. return UserAPIKeyAuth() - request_data = await _read_request_body(request=request) + request_data = await read_request_body( # rebind-ok: pre-existing rebinding on a rename-only line + request=request + ) request_data = populate_request_with_path_params(request_data=request_data, request=request) - return await _user_api_key_auth_builder( + return await user_api_key_auth_builder( request=request, api_key=api_key, azure_api_key_header="", @@ -2275,7 +2285,7 @@ if MCP_AVAILABLE: authorized: Final = await catalog.resolve( server_id, user_api_key_dict, - is_admin_view=_user_has_admin_view(user_api_key_dict), + is_admin_view=user_api_key_has_admin_view(user_api_key_dict), not_found_detail={"error": f"MCP server {server_id} not found"}, forbidden_detail={"error": f"Access denied to MCP server {server_id}"}, non_admin_missing="not_found", @@ -2313,7 +2323,7 @@ if MCP_AVAILABLE: scope: str | None = None, ): async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server: - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) # Use the server's stored client_id when the caller doesn't supply one stored_or_supplied_client_id: Final = mcp_server.client_id or client_id or "" ephemeral_dcr_client: Final = ( @@ -2373,7 +2383,7 @@ if MCP_AVAILABLE: scope: str | None = Form(None), ): async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server: - _raise_if_not_oauth2(mcp_server) + raise_if_not_oauth2(mcp_server) # Sealed passthrough codes exist only for the authorization_code grant. A refresh_token # grant must never open one: the minted client is unrecoverable after the single flow by # contract, so an expired browser-held token re-runs authorize instead. @@ -2425,7 +2435,7 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server: - request_data: Final = await _read_request_body(request=request) + request_data: Final = await read_request_body(request=request) data: Final[Mapping[str, object]] = {**request_data} client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris")) client_application_type: Final = client_supplied_application_type(data.get("application_type")) @@ -2748,7 +2758,7 @@ if MCP_AVAILABLE: servers: Final = {srv.server_id: srv for srv in await get_mcp_servers(prisma_client, server_ids)} allowed_server_ids: Final = ( None - if _user_has_admin_view(user_api_key_dict) + if user_api_key_has_admin_view(user_api_key_dict) else frozenset[str]().union( *[ await global_mcp_server_manager.get_allowed_mcp_servers(context) @@ -2839,7 +2849,7 @@ if MCP_AVAILABLE: authorized: Final = await catalog.resolve( server_id, user_api_key_dict, - is_admin_view=_user_has_admin_view(user_api_key_dict), + is_admin_view=user_api_key_has_admin_view(user_api_key_dict), not_found_detail={"error": f"MCP Server {server_id} not found"}, forbidden_detail={ "error": ( @@ -3317,7 +3327,7 @@ if MCP_AVAILABLE: Used by the UI to show a discovery grid when adding new MCP servers. """ # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={ @@ -3373,7 +3383,7 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={ @@ -3452,7 +3462,7 @@ if MCP_AVAILABLE: ): """Return toolsets the calling key is allowed to access.""" prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): op: Final = user_api_key_dict.object_permission if op is None or not op.mcp_toolsets: return await list_mcp_toolsets(prisma_client) @@ -3472,7 +3482,7 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - if not _user_has_admin_view(user_api_key_dict) and toolset_id not in await granted_toolset_ids( + if not user_api_key_has_admin_view(user_api_key_dict) and toolset_id not in await granted_toolset_ids( user_api_key_dict ): raise HTTPException( diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index c7883f5cf56..deefea6259b 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -77,9 +77,10 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient from litellm.proxy.management.teams.authz import TEAM_ADMIN_ONLY, is_team_admin from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.team_endpoints import ( - _refresh_cached_team, +from litellm.proxy.management_endpoints.team_endpoints import ( # noqa: F401 # legacy module exports + _refresh_cached_team, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export append_team_models, + refresh_cached_team, team_model_add, team_model_delete, ) @@ -1548,7 +1549,7 @@ async def unblock_model( #################################################################################### -async def _add_model_to_db( +async def add_model_to_db( model_params: Deployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, @@ -1591,7 +1592,10 @@ async def _add_model_to_db( return await table.create(data=_create_data) -async def _add_team_model_to_db( +_add_model_to_db: Final = add_model_to_db + + +async def add_team_model_to_db( model_params: Deployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, @@ -1625,7 +1629,7 @@ async def _add_team_model_to_db( model_params.model_name = unique_model_name ## CREATE MODEL IN DB ## - model_response: Final = await _add_model_to_db( + model_response: Final = await add_model_to_db( model_params=model_params, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1646,6 +1650,9 @@ async def _add_team_model_to_db( return model_response +_add_team_model_to_db: Final = add_team_model_to_db + + async def _update_team_model_in_db( db_model: Deployment, patch_data: updateDeployment, @@ -1966,7 +1973,7 @@ async def _remove_unbacked_team_models( data={"models": [model for model in existing_team_row.models if model not in names_to_remove]}, include={"object_permission": True}, ) - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team_row, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -2600,9 +2607,7 @@ async def add_new_model( reload_outcome: ReconcileOutcome = ReconcileOutcome(still_desired=None, live_after=None) try: _original_litellm_model_name: Final = model_params.model_name - add_model: Final = ( - _add_model_to_db if model_params.model_info.team_id is None else _add_team_model_to_db - ) + add_model: Final = add_model_to_db if model_params.model_info.team_id is None else add_team_model_to_db model_response = await add_model( model_params=priced_model_params, user_api_key_dict=user_api_key_dict, @@ -3202,7 +3207,7 @@ async def get_auto_router_classifier_default_prompt( ) -def _deduplicate_litellm_router_models(models: list[dict]) -> list[dict]: +def deduplicate_litellm_router_models(models: list[dict]) -> list[dict]: """ Deduplicate models based on their model_info.id field. Returns a list of unique models keeping only the first occurrence of each model ID. @@ -3223,6 +3228,9 @@ def _deduplicate_litellm_router_models(models: list[dict]) -> list[dict]: return unique_models +_deduplicate_litellm_router_models: Final = deduplicate_litellm_router_models + + _JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) @@ -3457,7 +3465,7 @@ async def clear_cache() -> ReconcileOutcome: # Reload only DB models. _add_deployment_locked, not add_deployment: this # coroutine already holds MODEL_RECONCILE_LOCK and asyncio.Lock is not # reentrant, so the public wrapper would deadlock against itself. - outcome: Final = await proxy_config._add_deployment_locked( + outcome: Final = await proxy_config.add_deployment_locked( prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj ) diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index e6705730e6f..b8ffec639c0 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -48,9 +48,11 @@ from litellm.proxy.management_endpoints.budget_management_endpoints import ( update_budget, ) from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity -from litellm.proxy.management_endpoints.common_utils import ( - _set_object_metadata_field, - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _set_object_metadata_field, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + set_object_metadata_field, + user_api_key_has_admin_view, validate_budget_duration, ) from litellm.proxy.management_helpers.object_permission_utils import ( @@ -262,7 +264,7 @@ def _table( return prisma_table -async def _verify_org_access( +async def verify_org_access( organization_id: str, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, @@ -272,7 +274,7 @@ async def _verify_org_access( Raises HTTPException(403) if the caller does not have access. """ - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): return if not user_api_key_dict.user_id: @@ -306,6 +308,9 @@ async def _verify_org_access( ) +_verify_org_access: Final = verify_org_access + + _STR_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) _BUDGET_SETTABLE_FIELDS: Final = frozenset(LiteLLM_BudgetTable.model_fields.keys()) - {"budget_id"} _ORG_COLUMN_FIELDS: Final = frozenset({"organization_alias", "models"}) @@ -537,7 +542,7 @@ async def new_organization( for field in _ORG_METADATA_FIELDS: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=organization_row, field_name=field, value=getattr(data, field), @@ -624,7 +629,7 @@ async def resolve_organization_daily_activity_scope( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, ) -> _OrganizationDailyActivityScope: - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) memberships: Final = ( await _table(OrganizationMembershipRepository(prisma_client)).find_many( where={"user_id": user_api_key_dict.user_id} @@ -750,7 +755,7 @@ async def update_organization( # IDOR guard: only proxy admins / org admins of THIS org may update # it. Without this, any authenticated key holder could rewrite # another organization's metadata, budgets, and object permissions. - await _verify_org_access( + await verify_org_access( organization_id=data.organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -927,7 +932,7 @@ async def update_organization_v2( }, ) - await _verify_org_access( + await verify_org_access( organization_id=organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1144,7 +1149,7 @@ async def list_organization( } # if proxy admin or admin viewer - get all orgs (with optional filters) - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): response = await _table(OrganizationRepository(prisma_client)).find_many( where=where_conditions if where_conditions else None, include={"litellm_budget_table": True, "members": True, "teams": True}, @@ -1210,7 +1215,7 @@ async def info_organization( raise HTTPException(status_code=500, detail={"error": "No db connected"}) # Verify caller has access to this organization - await _verify_org_access( + await verify_org_access( organization_id=organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1263,7 +1268,7 @@ async def deprecated_info_organization( # Verify caller has access to each requested organization for org_id in data.organizations: - await _verify_org_access( + await verify_org_access( organization_id=org_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1339,7 +1344,7 @@ async def organization_member_add( # organization, allowed to access this endpoint" — but the code # never enforced that. Any authenticated key holder could add # members to any org. Now gated explicitly. - await _verify_org_access( + await verify_org_access( organization_id=data.organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1453,7 +1458,7 @@ async def organization_member_update( # update member roles. The PROXY_ADMIN-target check below was # the only access control; without this, any authenticated user # could change any non-admin member's role in any org. - await _verify_org_access( + await verify_org_access( organization_id=data.organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -1602,7 +1607,7 @@ async def organization_member_delete( # IDOR guard: only proxy admins / org admins of THIS org may # delete members. Without this, any authenticated key holder # could remove any user from any org. - await _verify_org_access( + await verify_org_access( organization_id=data.organization_id, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, diff --git a/litellm/proxy/management_endpoints/policy_endpoints/__init__.py b/litellm/proxy/management_endpoints/policy_endpoints/__init__.py index 862c92bace9..6ba9872f322 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/__init__.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/__init__.py @@ -9,12 +9,20 @@ are imported directly into this namespace. from litellm.proxy.management_endpoints.policy_endpoints.endpoints import * # noqa: F403 from litellm.proxy.management_endpoints.policy_endpoints.endpoints import ( # noqa: F401 - _build_all_names_per_competitor, - _build_comparison_blocked_words, - _build_competitor_guardrail_definitions, - _build_name_blocked_words, - _build_recommendation_blocked_words, - _build_refinement_prompt, - _clean_competitor_line, - _parse_variations_response, + _build_all_names_per_competitor, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_comparison_blocked_words, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_competitor_guardrail_definitions, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_name_blocked_words, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_recommendation_blocked_words, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _build_refinement_prompt, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _clean_competitor_line, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + _parse_variations_response, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export + build_all_names_per_competitor, + build_comparison_blocked_words, + build_competitor_guardrail_definitions, + build_name_blocked_words, + build_recommendation_blocked_words, + build_refinement_prompt, + clean_competitor_line, + parse_variations_response, ) diff --git a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py index c3d1020f089..3bef7c107eb 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py @@ -757,7 +757,7 @@ async def enrich_policy_template( variations_map: Final = await _generate_competitor_variations(competitors, model=model) - enriched_definitions: Final = _build_competitor_guardrail_definitions( + enriched_definitions: Final = build_competitor_guardrail_definitions( template.get("guardrailDefinitions", []), competitors, brand_name, @@ -771,7 +771,7 @@ async def enrich_policy_template( } -def _build_refinement_prompt( +def build_refinement_prompt( instruction: str, existing_competitors: list[str], brand_name: str, @@ -788,6 +788,9 @@ def _build_refinement_prompt( ) +_build_refinement_prompt: Final = build_refinement_prompt + + async def _stream_llm_competitor_names( prompt: str, model: str, @@ -817,13 +820,13 @@ async def _stream_llm_competitor_names( buffer += delta while "\n" in buffer: line, buffer = buffer.split("\n", 1) - name = _clean_competitor_line(line) + name = clean_competitor_line(line) if name and name.lower() not in existing_lower and count < MAX_COMPETITOR_NAMES: existing_lower.add(name.lower()) count += 1 yield name, False # Handle remaining buffer - name = _clean_competitor_line(buffer) + name = clean_competitor_line(buffer) # rebind-ok: pre-existing rebinding on a rename-only line if name and name.lower() not in existing_lower and count < MAX_COMPETITOR_NAMES: yield name, False @@ -843,7 +846,7 @@ async def _stream_competitor_events( for comp in competitors: yield f"data: {json.dumps({'type': 'competitor', 'name': comp})}\n\n" - refinement_prompt: Final = _build_refinement_prompt(data.instruction, competitors, brand_name) + refinement_prompt: Final = build_refinement_prompt(data.instruction, competitors, brand_name) try: async for name, _ in _stream_llm_competitor_names(refinement_prompt, model, competitors): if name: @@ -875,7 +878,7 @@ async def _stream_competitor_events( total_variations: Final = sum(len(v) for v in variations_map.values()) yield f"data: {json.dumps({'type': 'status', 'message': f'Building guardrail definitions with {total_variations} variations...'})}\n\n" - enriched_definitions: Final = _build_competitor_guardrail_definitions( + enriched_definitions: Final = build_competitor_guardrail_definitions( template.get("guardrailDefinitions", []), competitors, brand_name, @@ -916,12 +919,15 @@ async def enrich_policy_template_stream( ) -def _clean_competitor_line(line: str) -> str | None: +def clean_competitor_line(line: str) -> str | None: """Strip numbering, bullets, and whitespace from a competitor name line.""" name: Final = line.strip().strip(".-) ").strip() return name if name and len(name) > 1 else None +_clean_competitor_line: Final = clean_competitor_line + + async def _generate_competitor_variations(competitors: list, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL) -> dict: """Generate common misspellings, abbreviations, and alternate names for each competitor.""" if not competitors: @@ -951,13 +957,13 @@ async def _generate_competitor_variations(competitors: list, model: str = DEFAUL temperature=COMPETITOR_LLM_TEMPERATURE, ) raw: Final = response.choices[0].message.content or "" - return _parse_variations_response(raw, capped) + return parse_variations_response(raw, capped) except Exception as e: verbose_proxy_logger.error("LLM competitor variation generation failed: %s", e) return {} -def _parse_variations_response(raw: str, competitors: list) -> dict[str, list[str]]: +def parse_variations_response(raw: str, competitors: list) -> dict[str, list[str]]: """Parse the LLM response for competitor variations into a name -> variations map.""" # Build a lowercase lookup for case-insensitive matching lower_to_canonical: Final = {comp.lower(): comp for comp in competitors} @@ -978,6 +984,9 @@ def _parse_variations_response(raw: str, competitors: list) -> dict[str, list[st return variations_map +_parse_variations_response: Final = parse_variations_response + + async def _discover_competitors_via_llm(prompt: str, model: str = DEFAULT_COMPETITOR_DISCOVERY_MODEL) -> list: """Call an onboarded LLM to discover competitor names.""" try: @@ -991,21 +1000,26 @@ async def _discover_competitors_via_llm(prompt: str, model: str = DEFAULT_COMPET temperature=COMPETITOR_LLM_TEMPERATURE, ) raw: Final = response.choices[0].message.content or "" - competitors = [name for line in raw.strip().split("\n") if (name := _clean_competitor_line(line)) is not None] + competitors: Final = [ + name for line in raw.strip().split("\n") if (name := clean_competitor_line(line)) is not None + ] return competitors[:MAX_COMPETITOR_NAMES] except Exception as e: verbose_proxy_logger.error("LLM competitor discovery failed: %s", e) return [] -def _build_all_names_per_competitor( +def build_all_names_per_competitor( competitors: list[str], variations_map: dict[str, list[str]] ) -> dict[str, list[str]]: """Build canonical + variation name lists for each competitor.""" return {comp: [comp] + variations_map.get(comp, []) for comp in competitors} -def _build_competitor_guardrail_definitions( +_build_all_names_per_competitor: Final = build_all_names_per_competitor + + +def build_competitor_guardrail_definitions( definitions: list, competitors: list, brand_name: str, @@ -1014,11 +1028,11 @@ def _build_competitor_guardrail_definitions( """Build enriched guardrailDefinitions with competitor names and variations populated.""" variations_map = variations_map or {} enriched: Final = copy.deepcopy(definitions) - all_names: Final = _build_all_names_per_competitor(competitors, variations_map) + all_names: Final = build_all_names_per_competitor(competitors, variations_map) - output_blocked: Final = _build_name_blocked_words(competitors, all_names) - recommendation_blocked: Final = _build_recommendation_blocked_words(competitors, all_names) - comparison_blocked: Final = _build_comparison_blocked_words(competitors, all_names, brand_name) + output_blocked: Final = build_name_blocked_words(competitors, all_names) + recommendation_blocked: Final = build_recommendation_blocked_words(competitors, all_names) + comparison_blocked: Final = build_comparison_blocked_words(competitors, all_names, brand_name) blocked_words_map: Final = { "competitor-output-blocker": output_blocked, @@ -1042,7 +1056,10 @@ def _build_competitor_guardrail_definitions( return enriched -def _build_name_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]: +_build_competitor_guardrail_definitions: Final = build_competitor_guardrail_definitions + + +def build_name_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]: """Build blocked word entries for direct competitor name mentions.""" result: Final = [] for comp in competitors: @@ -1052,7 +1069,10 @@ def _build_name_blocked_words(competitors: list[str], all_names: dict[str, list[ return result -def _build_recommendation_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]: +_build_name_blocked_words: Final = build_name_blocked_words + + +def build_recommendation_blocked_words(competitors: list[str], all_names: dict[str, list[str]]) -> list[dict]: """Build blocked word entries for competitor recommendations.""" result: Final = [] for comp in competitors: @@ -1068,7 +1088,10 @@ def _build_recommendation_blocked_words(competitors: list[str], all_names: dict[ return result -def _build_comparison_blocked_words( +_build_recommendation_blocked_words: Final = build_recommendation_blocked_words + + +def build_comparison_blocked_words( competitors: list[str], all_names: dict[str, list[str]], brand_name: str ) -> list[dict]: """Build blocked word entries for unfavorable competitor comparisons.""" @@ -1102,6 +1125,9 @@ def _build_comparison_blocked_words( return result +_build_comparison_blocked_words: Final = build_comparison_blocked_words + + class SuggestTemplatesRequest(LiteLLMBaseModel): attack_examples: list[str] = Field(default_factory=list) description: str = Field(default="") diff --git a/litellm/proxy/management_endpoints/prompt_cache_prediction.py b/litellm/proxy/management_endpoints/prompt_cache_prediction.py index 9103fa09893..a469f1d4f8d 100644 --- a/litellm/proxy/management_endpoints/prompt_cache_prediction.py +++ b/litellm/proxy/management_endpoints/prompt_cache_prediction.py @@ -15,12 +15,14 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import can_key_call_resolved_model from litellm.proxy.auth.auth_utils import get_cache_prediction_deployments from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # canonical parsed-body owner; validate its legacy result at the endpoint boundary +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # canonical parsed-body owner; validate its legacy result at the endpoint boundary ) from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms, predict_arm -from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( # noqa: F401 # legacy module exports + PROXY_MaxParallelRequestsHandler_v3, + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.types.llms.base import LiteLLMBaseModel @@ -39,7 +41,7 @@ class _CallerSettings(LiteLLMBaseModel): def _capacity_counter( - limiter: _PROXY_MaxParallelRequestsHandler_v3, + limiter: PROXY_MaxParallelRequestsHandler_v3, caller: UserAPIKeyAuth, model_name: str, request_data: Mapping[str, object], @@ -114,7 +116,7 @@ async def predict_cache_cost( or not caller or unsupported_transform or unsupported_headers - or not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3) + or not isinstance(limiter, PROXY_MaxParallelRequestsHandler_v3) ): reason: Final = ( "unsupported_provider_headers" @@ -134,7 +136,7 @@ async def predict_cache_cost( cache_rebuild_penalty=None, ) request_data: Final = _capacity_request_data( - http_request, user_api_key_dict, _REQUEST_DATA.validate_python(await _read_request_body(http_request)) + http_request, user_api_key_dict, _REQUEST_DATA.validate_python(await read_request_body(http_request)) ) stay: Final = await predict_arm( current, diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 12e1a0841f4..04886072856 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -45,10 +45,16 @@ from litellm.proxy._types import ( TeamMemberDeleteRequest, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_checks import _delete_cache_key_object +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _delete_cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + delete_cache_key_object, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) from litellm.proxy.management_endpoints.internal_user_endpoints import new_user from litellm.proxy.management_endpoints.scim.scim_transformations import ( ScimTransformations, @@ -58,10 +64,11 @@ from litellm.proxy.management_endpoints.team_endpoints import ( team_member_add, team_member_delete, ) -from litellm.proxy.utils import ( +from litellm.proxy.utils import ( # noqa: F401 # legacy module exports PrismaClient, - _premium_user_check, + _premium_user_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export handle_exception_on_proxy, + premium_user_check, ) from litellm.repositories.table_repositories import ( InvitationLinkRepository, @@ -264,7 +271,7 @@ class GroupMemberExtractionResult(LiteLLMBaseModel): scim_router: Final = APIRouter( prefix="/scim/v2", tags=["✨ SCIM v2 (Enterprise Only)"], - dependencies=[Depends(_premium_user_check)], + dependencies=[Depends(premium_user_check)], ) SCIM_MAX_PAGE_SIZE: Final = 100 @@ -1045,7 +1052,7 @@ async def _set_user_keys_blocked(user_id: str, blocked: bool) -> int: ) for key_row in affected_keys: - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=key_row.token, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -1546,7 +1553,7 @@ async def get_service_provider_config(request: Request): "SCIM ServiceProviderConfig request: method=%s url=%s headers=%s", request.method, request.url, - _safe_get_request_headers(request), + safe_get_request_headers(request), ) meta: Final = { "resourceType": "ServiceProviderConfig", diff --git a/litellm/proxy/management_endpoints/session_endpoints.py b/litellm/proxy/management_endpoints/session_endpoints.py index 2ba84bf03e5..4a4b81f7844 100644 --- a/litellm/proxy/management_endpoints/session_endpoints.py +++ b/litellm/proxy/management_endpoints/session_endpoints.py @@ -30,8 +30,9 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.auth_checks import delete_cache_key_objects from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _persist_deleted_verification_tokens, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + persist_deleted_verification_tokens, ) from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, @@ -86,7 +87,7 @@ async def revoke_ui_session_keys( return 0 revoked_tokens: Final = _TOKEN_LIST.validate_python(tuple(row.token for row in revoked_rows)) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=revoked_rows, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -154,7 +155,7 @@ async def session_logout( caller_row: Final = cast( # cast-ok: find_unique returns a prisma row shaped like the pydantic model "LiteLLM_VerificationToken", row ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=(caller_row,), prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/management_endpoints/team_callback_endpoints.py b/litellm/proxy/management_endpoints/team_callback_endpoints.py index 471a0814ae5..91fd447fb8d 100644 --- a/litellm/proxy/management_endpoints/team_callback_endpoints.py +++ b/litellm/proxy/management_endpoints/team_callback_endpoints.py @@ -34,19 +34,24 @@ from litellm.proxy.common_utils.callback_config_validation import ( conflicting_span_scope_error, cross_entry_family_error, ) -from litellm.proxy.common_utils.callback_utils import ( - _CALLBACK_VAR_ENCRYPTED_PREFIX, +from litellm.proxy.common_utils.callback_utils import ( # noqa: F401 # legacy module exports + _CALLBACK_VAR_ENCRYPTED_PREFIX, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + CALLBACK_VAR_ENCRYPTED_PREFIX, decrypt_callback_vars, encrypt_callback_vars, is_sensitive_callback_key, ) -from litellm.proxy.litellm_pre_call_utils import ( - _get_validated_callback_metadata, +from litellm.proxy.litellm_pre_call_utils import ( # noqa: F401 # legacy module exports + _get_validated_callback_metadata, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export convert_key_logging_metadata_to_callback, + get_validated_callback_metadata, ) from litellm.proxy.management.teams.authz import TEAM_OR_ORG_ADMIN, team_access_denied from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.team_endpoints import _refresh_cached_team +from litellm.proxy.management_endpoints.team_endpoints import ( # noqa: F401 # legacy module exports + _refresh_cached_team, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + refresh_cached_team, +) from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.repositories.team_repository import TeamRepository @@ -114,7 +119,7 @@ def _mask_sensitive_callback_vars(callbacks: TeamCallbackMetadata) -> None: return for key in tuple(callbacks.callback_vars): value = callbacks.callback_vars[key] - if is_sensitive_callback_key(key) or str(value).startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX): + if is_sensitive_callback_key(key) or str(value).startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): callbacks.callback_vars[key] = _CALLBACK_VARS_REDACTED @@ -147,7 +152,7 @@ def _resolve_team_callbacks(team_metadata: object) -> TeamCallbackMetadata: for entry in logging_entries if isinstance(logging_entries, list) else (): if not isinstance(entry, dict): continue - callback = _get_validated_callback_metadata(item=entry, source="team-level read") + callback = get_validated_callback_metadata(item=entry, source="team-level read") if callback is None: continue resolved = convert_key_logging_metadata_to_callback(data=callback, team_callback_settings_obj=resolved) @@ -399,7 +404,7 @@ async def add_team_callbacks( raise _callback_error(400, f"Team id = {team_id} does not exist. Please use a different team id.") # Without this a newly registered callback stays dormant for existing keys. - await _refresh_cached_team( + await refresh_cached_team( team_row=new_team_row, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -529,7 +534,7 @@ async def delete_team_callback( # Request-time callback resolution reads the cached team, so without this # the removed callback keeps firing for live keys until the cache expires. - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -671,7 +676,7 @@ async def disable_team_logging( # Request-time callback resolution reads the cached team, so without this # the DB says logging is off while live keys keep sending until it expires. - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index dfd3b93ff23..e9e3a4dfc48 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -94,10 +94,11 @@ from litellm.proxy._types import ( UpdateTeamRequest, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_checks import ( +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports OrganizationNotFoundError, - _cache_team_object, + _cache_team_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export allowed_route_check_inside_route, + cache_team_object, can_org_access_model, delete_cache_key_objects, delete_cache_team_object, @@ -129,15 +130,22 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( InvalidDateRange, parse_canonical_date_range, ) -from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, - _check_passthrough_routes_caller_permission, - _set_object_metadata_field, - _team_member_has_permission, - _update_metadata_fields, - _upsert_budget_and_membership, - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _check_disable_global_guardrails_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _check_passthrough_routes_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _set_object_metadata_field, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _team_member_has_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _update_metadata_fields, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + check_disable_global_guardrails_caller_permission, + check_passthrough_routes_caller_permission, member_budget_patch, + set_object_metadata_field, + team_member_has_permission, + update_metadata_fields, + upsert_budget_and_membership, + user_api_key_has_admin_view, validate_budget_duration, validate_team_model_max_budget, ) @@ -161,11 +169,12 @@ from litellm.proxy.management_helpers.access_group_team_sync import ( reconcile_team_access_group_membership, sync_team_access_group_membership, ) -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export enforce_all_proxy_mcp_servers_grant_is_admin_only, handle_update_object_permission_common, invalidate_cached_object_permissions, + set_object_permission, ) from litellm.proxy.management_helpers.team_member_permission_checks import ( TeamMemberPermissionChecks, @@ -454,7 +463,7 @@ def _sanitize_for_log(value: object) -> str: return text.replace("\r", "").replace("\n", "") -async def _refresh_cached_team( +async def refresh_cached_team( team_row: _CacheableTeamRow, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, @@ -473,7 +482,7 @@ async def _refresh_cached_team( via `model_dump()` to match the cache shape `_cache_team_object` expects. """ - await _cache_team_object( + await cache_team_object( team_id=team_row.team_id, team_table=LiteLLM_TeamTableCachedObj.model_validate(team_row.model_dump()), user_api_key_cache=user_api_key_cache, @@ -481,6 +490,9 @@ async def _refresh_cached_team( ) +_refresh_cached_team: Final = refresh_cached_team + + _GENERAL_SETTINGS: Final = TypeAdapter(dict[str, object]) @@ -596,7 +608,7 @@ class TeamMemberBudgetHandler: new_team_data_json["metadata"]["team_member_budget_id"] = team_member_budget_table.budget_id # Remove team member fields from new_team_data_json - TeamMemberBudgetHandler._clean_team_member_fields(new_team_data_json) + TeamMemberBudgetHandler.clean_team_member_fields(new_team_data_json) return new_team_data_json @@ -668,17 +680,19 @@ class TeamMemberBudgetHandler: ) # Remove team member fields from updated_kv - TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) + TeamMemberBudgetHandler.clean_team_member_fields(updated_kv) return updated_kv @staticmethod - def _clean_team_member_fields(data_dict: dict) -> None: + def clean_team_member_fields(data_dict: dict) -> None: """Remove team member fields from data dictionary""" data_dict.pop("team_member_budget", None) data_dict.pop("team_member_budget_duration", None) data_dict.pop("team_member_rpm_limit", None) data_dict.pop("team_member_tpm_limit", None) + _clean_team_member_fields = clean_team_member_fields + @staticmethod async def clear_team_member_budget_fields( team_table: _TeamBudgetRow, @@ -712,7 +726,7 @@ class TeamMemberBudgetHandler: user_api_key_dict=user_api_key_dict, ) - TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) + TeamMemberBudgetHandler.clean_team_member_fields(updated_kv) return updated_kv @staticmethod @@ -1603,8 +1617,8 @@ async def new_team( if not creating_user_in_list: data.members_with_roles.append(Member(role="admin", user_id=user_api_key_dict.user_id)) - _check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team") - _check_disable_global_guardrails_caller_permission( + check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team") + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -1651,7 +1665,7 @@ async def new_team( is_proxy_admin=user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN, prisma_client=prisma_client, ) - data_json = await _set_object_permission( + data_json = await set_object_permission( # rebind-ok: pre-existing rebinding on a rename-only line data_json=data_json, prisma_client=prisma_client, ) @@ -1682,7 +1696,7 @@ async def new_team( # Set Management Endpoint Metadata Fields for field in LiteLLM_ManagementEndpoint_MetadataFields_Premium: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=complete_team_data, field_name=field, value=getattr(data, field), @@ -1690,7 +1704,7 @@ async def new_team( for field in LiteLLM_ManagementEndpoint_MetadataFields: if getattr(data, field, None) is not None: - _set_object_metadata_field( + set_object_metadata_field( object_data=complete_team_data, field_name=field, value=getattr(data, field), @@ -2285,13 +2299,13 @@ async def update_team( entity="team", ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( data, user_api_key_dict, entity="team", existing_metadata=_existing_team_metadata if isinstance(_existing_team_metadata, dict) else None, ) - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( data.disable_global_guardrails, data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict user_api_key_dict, @@ -2505,7 +2519,7 @@ async def update_team( explicitly_set_fields=_team_member_fields_in_request, ) else: - TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) + TeamMemberBudgetHandler.clean_team_member_fields(updated_kv) # Check object permission if data.object_permission is not None: @@ -2521,7 +2535,7 @@ async def update_team( ) # update team metadata fields - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) if updated_kv.get("metadata") is not None: updated_kv["metadata"] = encrypt_callback_vars(updated_kv["metadata"]) @@ -2558,7 +2572,7 @@ async def update_team( object_permission_ids=(existing_team.object_permission_id, team_row.object_permission_id), user_api_key_cache=user_api_key_cache, ) - await _refresh_cached_team( + await refresh_cached_team( team_row=team_row, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -3525,7 +3539,7 @@ def _is_member_addressed_by(member: Member, data: TeamMemberDeleteRequest) -> bo ) -def _cleanup_members_with_roles( +def cleanup_members_with_roles( existing_team_row: LiteLLM_TeamTable, data: TeamMemberDeleteRequest, ) -> tuple[tuple[Member, ...], list[Member]]: @@ -3542,6 +3556,9 @@ def _cleanup_members_with_roles( return removed_team_members, new_team_members +_cleanup_members_with_roles: Final = cleanup_members_with_roles + + @router.post( "/team/member_delete", tags=["team management"], @@ -3651,7 +3668,7 @@ async def _team_member_delete( detail={"error": f"Team id={data.team_id} does not exist in db"}, ) - removed_team_members, new_team_members = _cleanup_members_with_roles( + removed_team_members, new_team_members = cleanup_members_with_roles( existing_team_row=LiteLLM_TeamTable(team_id=data.team_id, members_with_roles=fresh_members), data=data, ) @@ -3717,10 +3734,10 @@ async def _team_member_delete( if user_ids_to_delete: if keys_to_delete: from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, + persist_deleted_verification_tokens, ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=keys_to_delete, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -3887,7 +3904,7 @@ async def team_member_update( if data.role is not None else None ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( tx=tx, team_id=data.team_id, user_id=received_user_id, @@ -4431,7 +4448,7 @@ async def delete_team( ## DELETE ASSOCIATED KEYS # Fetch keys before deletion to persist them from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, + persist_deleted_verification_tokens, ) keys_to_delete: Final = await _tokens_db(prisma_client).find_many(where={"team_id": {"in": data.team_ids}}) @@ -4441,7 +4458,7 @@ async def delete_team( ) if keys_to_delete: - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=keys_to_delete, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -5545,7 +5562,7 @@ async def _enforce_list_team_v2_access( Returns the (possibly overridden) user_id, org_admin_org_ids and, for an org admin's own query, the caller's own team ids. """ - is_proxy_admin: Final = _user_has_admin_view(user_api_key_dict) + is_proxy_admin: Final = user_api_key_has_admin_view(user_api_key_dict) caller_user_id: Final = user_api_key_dict.user_id if is_proxy_admin: @@ -5809,7 +5826,7 @@ async def _authorize_and_filter_teams( - Own query (user_id matches caller): teams the user is a member of, across all orgs. - Others: 401. """ - is_proxy_admin: Final = _user_has_admin_view(user_api_key_dict) + is_proxy_admin: Final = user_api_key_has_admin_view(user_api_key_dict) is_own_query: Final = ( user_id is not None and user_api_key_dict.user_id is not None and user_api_key_dict.user_id == user_id ) @@ -6169,7 +6186,7 @@ async def append_team_models( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -6252,7 +6269,7 @@ async def team_model_delete( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - await _refresh_cached_team( + await refresh_cached_team( team_row=updated_team, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, @@ -6300,7 +6317,7 @@ async def team_member_permissions( # a Proxy Admin would. Team / org admins keep their existing scope. if ( hasattr(user_api_key_dict, "user_role") - and not _user_has_admin_view(user_api_key_dict) + and not user_api_key_has_admin_view(user_api_key_dict) and not await get_team_access().allows(user_api_key_dict, complete_team_data, TEAM_OR_ORG_ADMIN) and not _is_available_team( team_id=complete_team_data.team_id, @@ -6556,7 +6573,7 @@ async def resolve_team_daily_activity_scope( if exclude_team_ids: exclude_team_ids_list = exclude_team_ids.split(",") if exclude_team_ids else None - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): user_info: Final = await get_user_object( user_id=user_api_key_dict.user_id, prisma_client=prisma_client, @@ -6598,12 +6615,12 @@ async def resolve_team_daily_activity_scope( # filtering the entire response by their own API keys (they can re- # request the admin-only teams separately to get the wider view). user_api_keys: list[str] | None = None - if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases: + if not user_api_key_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases: has_full_team_view = True for team_alias in team_aliases: team_obj = LiteLLM_TeamTable.model_validate(team_alias.model_dump()) is_admin = is_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj) - has_perm = _team_member_has_permission( + has_perm = team_member_has_permission( user_api_key_dict=user_api_key_dict, team_obj=team_obj, permission="/team/daily/activity", diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 605952b44f0..c515108ab11 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -90,8 +90,9 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object -from litellm.proxy.auth.auth_utils import ( - _get_request_ip_address, +from litellm.proxy.auth.auth_utils import ( # noqa: F401 # legacy module exports + _get_request_ip_address, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_request_ip_address, has_user_setup_sso, ) from litellm.proxy.auth.handle_jwt import JWTHandler @@ -318,7 +319,7 @@ def _cli_sso_start_response_body( def _get_cli_sso_start_rate_limit_cache_key(request: Request, use_x_forwarded_for: bool | None = False) -> str: - client_ip: Final = _get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) or "unknown" + client_ip: Final = get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) or "unknown" client_ip_hash: Final = _hash_cli_sso_secret(client_ip) return f"{_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX}:{client_ip_hash}" @@ -1060,7 +1061,7 @@ async def google_login( _get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache) # Store CLI login handle in state for OAuth flow - cli_state: Final[str | None] = SSOAuthenticationHandler._get_cli_state( + cli_state: Final[str | None] = SSOAuthenticationHandler.get_cli_state( source=source, key=key, user_code=(user_code if _cli_sso_verification_uri_complete_enabled() else None), @@ -1626,7 +1627,7 @@ async def get_generic_sso_response( param="code", code=status.HTTP_400_BAD_REQUEST, ) - combined_response: Final = await SSOAuthenticationHandler._pkce_token_exchange( + combined_response: Final = await SSOAuthenticationHandler.pkce_token_exchange( authorization_code=authorization_code, code_verifier=code_verifier, client_id=generic_client_id, @@ -1664,7 +1665,7 @@ async def get_generic_sso_response( # successfully. Deleting earlier would consume the verifier on a transient # failure, forcing the user to restart the entire OAuth flow from scratch. if pkce_cache_key: - await SSOAuthenticationHandler._delete_pkce_verifier(pkce_cache_key) + await SSOAuthenticationHandler.delete_pkce_verifier(pkce_cache_key) except Exception as e: _handle_generic_sso_error( @@ -2214,7 +2215,7 @@ async def saml_callback(request: Request): relay_state: Final = post_data.get("RelayState") cp_return_to: Final[str | None] = ( relay_state - if isinstance(relay_state, str) and SSOAuthenticationHandler._validate_return_to(relay_state) + if isinstance(relay_state, str) and SSOAuthenticationHandler.validate_return_to(relay_state) else None ) @@ -2433,7 +2434,7 @@ async def cli_sso_callback( result_non_none: Final[OpenID | dict] = cast(OpenID | dict, result) try: - parsed_openid_result: Final = SSOAuthenticationHandler._get_user_email_and_id_from_result( + parsed_openid_result: Final = SSOAuthenticationHandler.get_user_email_and_id_from_result( result=result_non_none, generic_client_id=os.getenv("GENERIC_CLIENT_ID", None), ) @@ -2808,7 +2809,7 @@ def _is_same_origin_return_path(return_to: str) -> bool: @with_service_target(SSO_SESSIONS_TARGET) -async def _sso_return_to_redirect( +async def sso_return_to_redirect( return_to: str | None, jwt_token: str, redis_usage_cache, @@ -2836,7 +2837,7 @@ async def _sso_return_to_redirect( redirect_response.delete_cookie("litellm_cp_return_to") return redirect_response - if SSOAuthenticationHandler._validate_return_to(return_to): + if SSOAuthenticationHandler.validate_return_to(return_to): code: Final = secrets.token_urlsafe(32) cache_key: Final = f"login_code:{code}" cache_value: Final = {"token": jwt_token, "redirect_url": return_to} @@ -2855,6 +2856,9 @@ async def _sso_return_to_redirect( return None +_sso_return_to_redirect: Final = sso_return_to_redirect + + def set_session_token_cookie(response: Response, request: Request, jwt_token: str) -> None: """Set the ``token`` session cookie shared by every sign-in path. @@ -2884,7 +2888,7 @@ def _persist_return_to_cookie(response: Response, return_to: str | None, request if return_to is None: return try: - safe: Final = _is_same_origin_return_path(return_to) or SSOAuthenticationHandler._validate_return_to(return_to) + safe: Final = _is_same_origin_return_path(return_to) or SSOAuthenticationHandler.validate_return_to(return_to) except HTTPException: return # a non-matching absolute return_to is ignored, never blocks sign-in if safe: @@ -2904,7 +2908,7 @@ class SSOAuthenticationHandler: """ @staticmethod - def _validate_return_to(return_to: str) -> bool: + def validate_return_to(return_to: str) -> bool: """ Validate that return_to matches the configured control_plane_url origin. @@ -2934,6 +2938,8 @@ class SSOAuthenticationHandler: return True + _validate_return_to = validate_return_to + @staticmethod async def get_sso_login_redirect( redirect_url: str, @@ -3451,7 +3457,7 @@ class SSOAuthenticationHandler: return team_request @staticmethod - def _get_cli_state( + def get_cli_state( source: str | None, key: str | None, existing_key: str | None = None, @@ -3477,8 +3483,10 @@ class SSOAuthenticationHandler: else: return None + _get_cli_state = get_cli_state + @staticmethod - def _get_user_email_and_id_from_result( + def get_user_email_and_id_from_result( result: OpenID | dict | None, generic_client_id: str | None = None, ) -> ParsedOpenIDResult: @@ -3535,6 +3543,8 @@ class SSOAuthenticationHandler: user_role=user_role, ) + _get_user_email_and_id_from_result = get_user_email_and_id_from_result + @staticmethod async def get_redirect_response_from_openid( result: OpenID | dict | CustomOpenID, @@ -3564,7 +3574,7 @@ class SSOAuthenticationHandler: prisma_client: Final = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy") # User is Authe'd in - generate key for the UI to access Proxy - parsed_openid_result: Final = SSOAuthenticationHandler._get_user_email_and_id_from_result( + parsed_openid_result: Final = SSOAuthenticationHandler.get_user_email_and_id_from_result( result=result, generic_client_id=generic_client_id ) user_email: Final = parsed_openid_result.get("user_email") @@ -3729,7 +3739,7 @@ class SSOAuthenticationHandler: # Post-SSO return_to handling (the same-origin DCR round-trip and the control-plane # cross-origin code exchange) lives in one shared helper so this method stays inside the # complexity budget. None falls through to the dashboard redirect below. - return_to_redirect: Final = await _sso_return_to_redirect( + return_to_redirect: Final = await sso_return_to_redirect( return_to=return_to, jwt_token=jwt_token, redis_usage_cache=redis_usage_cache, @@ -3862,7 +3872,7 @@ class SSOAuthenticationHandler: strict_cache_miss: Final = os.getenv("PKCE_STRICT_CACHE_MISS", "false").lower() == "true" if strict_cache_miss: if empty_value_in_dict: - await SSOAuthenticationHandler._delete_pkce_verifier(cache_key) + await SSOAuthenticationHandler.delete_pkce_verifier(cache_key) raise ProxyException( message=( f"PKCE verifier for state '{state}' was found in cache but " @@ -3873,7 +3883,7 @@ class SSOAuthenticationHandler: code=status.HTTP_401_UNAUTHORIZED, ) elif cached_data is not None: - await SSOAuthenticationHandler._delete_pkce_verifier(cache_key) + await SSOAuthenticationHandler.delete_pkce_verifier(cache_key) verbose_proxy_logger.error( "PKCE verifier for state '%s' has an unrecognized format (type=%s); " "treating as a cache miss. Investigate the cached value — it may be " @@ -3916,7 +3926,7 @@ class SSOAuthenticationHandler: ) else: if cached_data is not None: - await SSOAuthenticationHandler._delete_pkce_verifier(cache_key) + await SSOAuthenticationHandler.delete_pkce_verifier(cache_key) verbose_proxy_logger.warning( "PKCE is enabled but verifier not found in cache for state '%s' " "(cache type: %s, raw data present: %s). " @@ -3928,7 +3938,7 @@ class SSOAuthenticationHandler: @staticmethod @with_service_target(SSO_SESSIONS_TARGET) - async def _delete_pkce_verifier(cache_key: str) -> None: + async def delete_pkce_verifier(cache_key: str) -> None: """Delete a single-use PKCE verifier from cache after a successful exchange. Failure is non-fatal: a leftover verifier is a minor security concern @@ -3948,6 +3958,8 @@ class SSOAuthenticationHandler: exc, ) + _delete_pkce_verifier = delete_pkce_verifier + @staticmethod def generate_pkce_params() -> tuple[str, str]: """ @@ -4032,7 +4044,7 @@ class SSOAuthenticationHandler: return token_response @staticmethod - async def _pkce_token_exchange( + async def pkce_token_exchange( authorization_code: str, code_verifier: str, client_id: str, @@ -4158,6 +4170,8 @@ class SSOAuthenticationHandler: # Case 3: field absent from token_response — leave userinfo value as-is. return merged + _pkce_token_exchange = pkce_token_exchange + @staticmethod async def _get_pkce_userinfo( access_token: str, diff --git a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py index f03259cb4e9..cff5bef05a0 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py @@ -47,14 +47,14 @@ async def usage_ai_chat( The AI agent has access to tools that query aggregated daily activity data. """ from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, require_caller_user_id_for_non_admin, + user_api_key_has_admin_view, ) from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( stream_usage_ai_chat, ) - is_admin: Final = _user_has_admin_view(user_api_key_dict) + is_admin: Final = user_api_key_has_admin_view(user_api_key_dict) if is_admin: user_id = user_api_key_dict.user_id else: diff --git a/litellm/proxy/management_helpers/access_group_key_sync.py b/litellm/proxy/management_helpers/access_group_key_sync.py index a5289b89cec..024967a34f6 100644 --- a/litellm/proxy/management_helpers/access_group_key_sync.py +++ b/litellm/proxy/management_helpers/access_group_key_sync.py @@ -33,8 +33,9 @@ from litellm.proxy._types import ( RegenerateKeyRequest, UpdateKeyRequest, ) -from litellm.proxy.auth.auth_checks import ( - _delete_cache_access_object, # pyright: ignore[reportPrivateUsage] # the access-group endpoints reach for this same cache primitive +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _delete_cache_access_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + delete_cache_access_object, # pyright: ignore[reportPrivateUsage] # the access-group endpoints reach for this same cache primitive ) from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper from litellm.repositories.table_repositories import AccessGroupRepository @@ -86,7 +87,7 @@ async def _invalidate_access_group_cache(access_group_id: str) -> None: """ from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache - await _delete_cache_access_object( + await delete_cache_access_object( access_group_id=access_group_id, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/management_helpers/access_group_team_sync.py b/litellm/proxy/management_helpers/access_group_team_sync.py index fbca95b9169..1205414b4fb 100644 --- a/litellm/proxy/management_helpers/access_group_team_sync.py +++ b/litellm/proxy/management_helpers/access_group_team_sync.py @@ -20,7 +20,10 @@ from typing import Final, Protocol from pydantic import TypeAdapter -from litellm.proxy.auth.auth_checks import _delete_cache_access_object +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _delete_cache_access_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + delete_cache_access_object, +) from litellm.proxy.db.db_span import db_span from litellm.types.llms.base import LiteLLMBaseModel @@ -99,7 +102,7 @@ async def invalidate_access_group_cache(access_group_id: str) -> None: """ from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache - await _delete_cache_access_object( + await delete_cache_access_object( access_group_id=access_group_id, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/management_helpers/auto_router_permissions.py b/litellm/proxy/management_helpers/auto_router_permissions.py index f2d0aaaec78..16f6b512e0d 100644 --- a/litellm/proxy/management_helpers/auto_router_permissions.py +++ b/litellm/proxy/management_helpers/auto_router_permissions.py @@ -19,12 +19,13 @@ from litellm.proxy._types import ( LitellmUserRoles, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_checks import ( - _check_team_member_model_access, # pyright: ignore[reportPrivateUsage] # shared membership authorization owner +from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports + _check_team_member_model_access, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export can_key_call_model, can_org_access_model, can_project_access_model, can_team_access_model, + check_team_member_model_access, # pyright: ignore[reportPrivateUsage] # shared membership authorization owner ) from litellm.proxy.auth.team_grants import team_model_aliases from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper @@ -222,7 +223,7 @@ async def authorize_member_auto_router_dependencies( llm_router=llm_router, prisma_client=prisma_client, ) - await _check_team_member_model_access( + await check_team_member_model_access( model=model, team_object=team, valid_token=scoped_actor, diff --git a/litellm/proxy/management_helpers/bulk_team_member_budgets.py b/litellm/proxy/management_helpers/bulk_team_member_budgets.py index 292cb637687..08169eab761 100644 --- a/litellm/proxy/management_helpers/bulk_team_member_budgets.py +++ b/litellm/proxy/management_helpers/bulk_team_member_budgets.py @@ -25,9 +25,10 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient from litellm.proxy.management.teams.authz import TEAM_OR_ORG_ADMIN from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.common_utils import ( - _upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage] # the single-member write, shared so the two surfaces cannot drift +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export member_budget_patch, + upsert_budget_and_membership, # pyright: ignore[reportPrivateUsage] # the single-member write, shared so the two surfaces cannot drift ) from litellm.proxy.management_helpers.audit_logs import create_object_audit_log from litellm.proxy.management_helpers.bulk_user_deletion import ( @@ -213,7 +214,7 @@ async def bulk_update_team_member_budgets( tx, frozenset(budget_id for budget_id in budget_id_of.values() if budget_id is not None) ) for index, user_id in applied: - await _upsert_budget_and_membership( + await upsert_budget_and_membership( tx=tx, team_id=team_id, user_id=user_id, diff --git a/litellm/proxy/management_helpers/bulk_user_creation.py b/litellm/proxy/management_helpers/bulk_user_creation.py index ecaec1b6260..cca37ebd7ed 100644 --- a/litellm/proxy/management_helpers/bulk_user_creation.py +++ b/litellm/proxy/management_helpers/bulk_user_creation.py @@ -42,15 +42,17 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import ( _update_internal_new_user_params, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # /user/new defaults; result validated below check_if_default_team_set, ) -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage] # same permission check /user/new uses +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage] # same permission check /user/new uses generate_key_helper_fn, # pyright: ignore[reportUnknownVariableType] # legacy untyped helper; result validated by _KEY_RESPONSE metadata_json_with_limits, ) from litellm.proxy.management_endpoints.organization_endpoints import organization_member_add from litellm.proxy.management_helpers.access_group_team_sync import TEAM_ADVISORY_LOCK_SQL -from litellm.proxy.management_helpers.object_permission_utils import ( - _set_object_permission, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared with /user/new; result validated below +from litellm.proxy.management_helpers.object_permission_utils import ( # noqa: F401 # legacy module exports + _set_object_permission, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + set_object_permission, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared with /user/new; result validated below ) from litellm.proxy.management_helpers.utils import ( _resolve_member_budget_id, # pyright: ignore[reportPrivateUsage] # shared with /team/member_add @@ -214,7 +216,7 @@ def _row_error(item: BulkNewUserItem, user_api_key_dict: UserAPIKeyAuth) -> str ) try: validate_budget_duration(item.budget_duration) - _check_permissions_caller_permission(data=item, user_api_key_dict=user_api_key_dict) + check_permissions_caller_permission(data=item, user_api_key_dict=user_api_key_dict) if item.auto_create_key: enforce_batch_limits_are_admin_only(item, None, user_api_key_dict, "key") except Exception as exc: # noqa: BLE001 # any validation failure is reported on this row only @@ -344,7 +346,7 @@ async def _prepare_user(user: _PendingUser, prisma_client: PrismaClient) -> _Pre data: Final = {**dumped, "user_id": user.user_id} data_json: Final = _JSON_OBJECT.validate_python(_update_internal_new_user_params(data, user.request)) with_permission: Final = _JSON_OBJECT.validate_python( - await _set_object_permission(data_json=data_json, prisma_client=prisma_client) + await set_object_permission(data_json=data_json, prisma_client=prisma_client) ) return _PreparedUser(user, _USER_ROW.validate_python(with_permission)) except Exception as exc: # noqa: BLE001 # any preparation failure is reported on this row only diff --git a/litellm/proxy/management_helpers/bulk_user_deletion.py b/litellm/proxy/management_helpers/bulk_user_deletion.py index b4e4afcbc2a..c92ebce8531 100644 --- a/litellm/proxy/management_helpers/bulk_user_deletion.py +++ b/litellm/proxy/management_helpers/bulk_user_deletion.py @@ -36,8 +36,9 @@ from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventH from litellm.proxy.list_api.common import PROBLEM_TYPE_BASE, ManagementProblem from litellm.proxy.management.teams.authz import TEAM_OR_ORG_ADMIN from litellm.proxy.management.teams.dependencies import get_team_access -from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, # pyright: ignore[reportPrivateUsage] # same audit path /key/delete uses +from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports + _persist_deleted_verification_tokens, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + persist_deleted_verification_tokens, # pyright: ignore[reportPrivateUsage] # same audit path /key/delete uses ) from litellm.proxy.management_helpers.access_group_team_sync import TEAM_ADVISORY_LOCK_SQL from litellm.proxy.utils import PrismaClient, ProxyLogging @@ -267,7 +268,7 @@ async def _remove_members_from_team( await _user_tx_db(tx).update(where=_eq_filter("user_id", row.user_id), data=teams_data) await _membership_tx_db(tx).delete_many(where=_team_users_filter(team_id, cleanup_ids)) if keys: - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=keys, # pyright: ignore[reportArgumentType] # generated row model carries the same columns as LiteLLM_VerificationToken prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -398,7 +399,7 @@ async def _delete_user_rows( prisma_client=prisma_client, ) if keys: - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=keys, # pyright: ignore[reportArgumentType] # generated row model carries the same columns as LiteLLM_VerificationToken prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index c389381dca4..86545e6e064 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -199,10 +199,10 @@ async def invalidate_cached_object_permissions( await evict_and_broadcast(cache_keys, user_api_key_cache) -async def _set_object_permission( - data_json: dict, +async def set_object_permission( + data_json: dict[str, object], prisma_client: PrismaClient | None, -): +) -> dict[str, object]: """ Creates the LiteLLM_ObjectPermissionTable record for the key/team. Handles permissions for vector stores and mcp servers. @@ -237,6 +237,9 @@ async def _set_object_permission( return data_json +_set_object_permission: Final = set_object_permission + + def _dedupe_preserving_order(values: list[str]) -> list[str]: seen: Final[set[str]] = set() result: Final[list[str]] = [] @@ -447,7 +450,7 @@ async def _resolve_team_allowed_mcp_servers( direct_servers: Final[list[str]] = team_object_permission.mcp_servers or [] if SpecialMCPServerName.all_proxy_servers.value in direct_servers: return _get_all_mcp_server_ids() - access_group_servers: Final[list[str]] = await MCPRequestHandler._get_mcp_servers_from_access_groups( + access_group_servers: Final[list[str]] = await MCPRequestHandler.get_mcp_servers_from_access_groups( team_object_permission.mcp_access_groups or [] ) raw_tool_perms = team_object_permission.mcp_tool_permissions or {} @@ -463,7 +466,7 @@ async def _resolve_team_allowed_mcp_servers( return _flatten_resolved_mcp_server_ids(resolved_servers) | unresolved_servers -def _get_allow_all_keys_server_ids() -> set[str]: +def get_allow_all_keys_server_ids() -> set[str]: """Return the set of MCP server IDs marked with allow_all_keys=True.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, @@ -472,6 +475,9 @@ def _get_allow_all_keys_server_ids() -> set[str]: return set(global_mcp_server_manager.get_allow_all_keys_server_ids()) +_get_allow_all_keys_server_ids: Final = get_allow_all_keys_server_ids + + def _get_all_mcp_server_ids() -> set[str]: """Return every MCP server id registered on the proxy (config + DB union).""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -558,7 +564,7 @@ async def _get_grandfathered_key_mcp_server_ids( ) -async def _get_team_allowed_mcp_servers( +async def get_team_allowed_mcp_servers( team_obj: Optional["LiteLLM_TeamTableCachedObj"], prisma_client: PrismaClient | None = None, ) -> set[str]: @@ -574,10 +580,10 @@ async def _get_team_allowed_mcp_servers( return set() from litellm.proxy.auth.auth_checks import ( - _get_mcp_server_ids_from_access_groups, # pyright: ignore[reportPrivateUsage] # same resolver runtime MCP auth calls + get_mcp_server_ids_from_access_groups, # pyright: ignore[reportPrivateUsage] # same resolver runtime MCP auth calls ) - access_group_servers: Final = await _get_mcp_server_ids_from_access_groups( + access_group_servers: Final = await get_mcp_server_ids_from_access_groups( access_group_ids=team_obj.access_group_ids or [], prisma_client=prisma_client, ) @@ -599,6 +605,9 @@ async def _get_team_allowed_mcp_servers( ) +_get_team_allowed_mcp_servers: Final = get_team_allowed_mcp_servers + + def _extract_requested_mcp_server_ids( object_permission: ObjectPermissionDict | None, ) -> set[str]: @@ -692,8 +701,8 @@ async def validate_key_mcp_servers_against_team( if not requested_servers and not requested_access_groups and not requested_toolsets: return object_permission - allow_all_keys_servers: Final = _get_allow_all_keys_server_ids() - team_allowed_servers: Final = await _get_team_allowed_mcp_servers( + allow_all_keys_servers: Final = get_allow_all_keys_server_ids() + team_allowed_servers: Final = await get_team_allowed_mcp_servers( team_obj=team_obj, prisma_client=prisma_client, ) diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index a076d8240c6..3e6d0e47265 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -65,7 +65,7 @@ class TeamMemberPermissionChecks: Main handler for checking if a team member can update a key """ from litellm.proxy.management_endpoints.key_management_endpoints import ( - _get_caller_team_role, + get_caller_team_role, ) # 1. Don't execute these checks if the user role is proxy admin @@ -85,7 +85,7 @@ class TeamMemberPermissionChecks: check_db_only=True, ) - caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) + caller_team_role: Final = get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) # 4. Check if the team member has permissions for the endpoint has_permission: Final = TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( @@ -152,7 +152,7 @@ class TeamMemberPermissionChecks: from fastapi import HTTPException from litellm.proxy.management_endpoints.key_management_endpoints import ( - _get_caller_team_role, + get_caller_team_role, ) # No-op when the request does not assign any access groups. @@ -173,7 +173,7 @@ class TeamMemberPermissionChecks: ), ) - caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) + caller_team_role: Final = get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) # Team admins always bypass (consistent with other member-permission checks). if caller_team_role == "admin": @@ -209,7 +209,7 @@ class TeamMemberPermissionChecks: Returns True if the user belongs to the team that the key is assigned to """ from litellm.proxy.management_endpoints.key_management_endpoints import ( - _get_caller_team_role, + get_caller_team_role, ) from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -223,7 +223,7 @@ class TeamMemberPermissionChecks: check_db_only=True, ) - caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) + caller_team_role: Final = get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) return caller_team_role is not None @staticmethod diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index b8af3950859..30aaaacad04 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -35,7 +35,10 @@ from litellm.proxy._types import ( # key request types; user request types; tea UserAPIKeyAuth, VirtualKeyEvent, ) -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET from litellm.proxy.utils import PrismaClient, jsonify_object @@ -713,7 +716,9 @@ async def _emit_management_endpoint_otel_span( ) route = get_request_route(http_request) - request_body: dict = await _read_request_body(request=http_request) + request_body: dict = await read_request_body( # rebind-ok: pre-existing rebinding on a rename-only line + request=http_request + ) else: route = func.__name__ request_body = {} diff --git a/litellm/proxy/ocr_endpoints/endpoints.py b/litellm/proxy/ocr_endpoints/endpoints.py index 81d9f8a7b43..94fe9354386 100644 --- a/litellm/proxy/ocr_endpoints/endpoints.py +++ b/litellm/proxy/ocr_endpoints/endpoints.py @@ -340,7 +340,7 @@ async def ocr( return _native_response(response, fastapi_response) or response except Exception as e: processor = ProxyBaseLLMRequestProcessing(data=data) - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/openai_evals_endpoints/endpoints.py b/litellm/proxy/openai_evals_endpoints/endpoints.py index abfbed5f822..aa1c977b8e7 100644 --- a/litellm/proxy/openai_evals_endpoints/endpoints.py +++ b/litellm/proxy/openai_evals_endpoints/endpoints.py @@ -107,7 +107,7 @@ async def create_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -208,7 +208,7 @@ async def list_evals( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -296,7 +296,7 @@ async def get_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -386,7 +386,7 @@ async def update_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -474,7 +474,7 @@ async def delete_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -562,7 +562,7 @@ async def cancel_eval( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -666,7 +666,7 @@ async def create_run( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -759,7 +759,7 @@ async def list_runs( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -846,7 +846,7 @@ async def get_run( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -935,7 +935,7 @@ async def cancel_run( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1024,7 +1024,7 @@ async def delete_run( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 7e7ab1dc1f0..ffeca8d0257 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: object) -> 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 @@ -121,8 +121,11 @@ def _is_base64_encoded_unified_file_id(b64_uid: object) -> str | Literal[False]: return False +_is_base64_encoded_unified_file_id: Final = is_base64_encoded_unified_file_id + + def convert_b64_uid_to_unified_uid(b64_uid: str) -> str: - is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(b64_uid) + is_base64_unified_file_id: Final = is_base64_encoded_unified_file_id(b64_uid) if is_base64_unified_file_id: return is_base64_unified_file_id else: @@ -942,10 +945,10 @@ async def extract_file_creation_params( Returns: FileCreationParams: Structured parameters extracted from the request """ - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body if request_body is None: - request_body = await _read_request_body(request=request) or {} + request_body = await read_request_body(request=request) or {} # Extract target_storage (simplified - just use form parameter) target_storage: Final = _extract_target_storage_simple(target_storage_form) @@ -1097,7 +1100,7 @@ async def validate_managed_id_requirement( if not resource_id: return - if not _is_base64_encoded_unified_file_id(resource_id): + if not is_base64_encoded_unified_file_id(resource_id): raise HTTPException( status_code=400, detail=( @@ -1150,7 +1153,7 @@ def _batch_response_model_id_candidates( ) -> tuple[str, ...]: response_id: Final = getattr(response, "id", None) decoded_response_id: Final = ( - _is_base64_encoded_unified_file_id(response_id) if isinstance(response_id, str) else False + is_base64_encoded_unified_file_id(response_id) if isinstance(response_id, str) else False ) return tuple( candidate @@ -1215,7 +1218,7 @@ async def resolve_input_file_id_to_unified(response, prisma_client) -> None: if ( hasattr(response, "input_file_id") and response.input_file_id - and not _is_base64_encoded_unified_file_id(response.input_file_id) + and not is_base64_encoded_unified_file_id(response.input_file_id) and prisma_client ): try: @@ -1238,7 +1241,7 @@ async def resolve_output_file_ids_to_unified(response, prisma_client) -> None: return for attr in ("output_file_id", "error_file_id"): raw_id = getattr(response, attr, None) - if not raw_id or _is_base64_encoded_unified_file_id(raw_id): + if not raw_id or is_base64_encoded_unified_file_id(raw_id): continue try: managed_file = await ManagedFileRepository(prisma_client).table.find_first( @@ -1307,7 +1310,7 @@ async def ensure_batch_response_managed_file_ids( for file_attr in ("output_file_id", "error_file_id"): raw_file_id = getattr(response, file_attr, None) - if not raw_file_id or _is_base64_encoded_unified_file_id(raw_file_id): + if not raw_file_id or is_base64_encoded_unified_file_id(raw_file_id): continue try: new_unified_file_id = managed_files_obj.get_unified_output_file_id( diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index e170e1cf894..4e31f9713d0 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -42,9 +42,10 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import ( resolve_litellm_executed_provider, ) from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export extract_nested_form_metadata, + read_request_body, ) from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_body, @@ -70,8 +71,8 @@ from litellm.proxy.openai_files_endpoints.batch_guardrails import ( rewrite_batch_input_file, scan_batch_input_file, ) -from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, +from litellm.proxy.openai_files_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export add_internal_model_credentials, apply_team_provider_credentials, authorize_model_for_key, @@ -79,6 +80,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( extract_file_creation_params, get_authorized_credentials_for_model, handle_model_based_routing, + is_base64_encoded_unified_file_id, prepare_data_with_credentials, validate_file_list_limit, validate_managed_files_requirement, @@ -620,7 +622,7 @@ async def create_file( ) # Extract file creation parameters using utility function - request_body: Final = await _read_request_body(request=request) or {} + request_body: Final = await read_request_body(request=request) or {} file_params: Final = await extract_file_creation_params( request=request, request_body=request_body, @@ -1019,7 +1021,7 @@ async def get_file_content( ) ## check if file_id is a litellm managed file - is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(file_id) + is_base64_unified_file_id: Final = is_base64_encoded_unified_file_id(file_id) if is_base64_unified_file_id: managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: @@ -1366,7 +1368,7 @@ async def get_file( ) ## EXISTING: check if file_id is a litellm managed file - elif _is_base64_encoded_unified_file_id(file_id): + elif is_base64_encoded_unified_file_id(file_id): managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: raise ProxyException( @@ -1576,7 +1578,7 @@ async def delete_file( ) ## EXISTING: check if file_id is a litellm managed file - elif _is_base64_encoded_unified_file_id(file_id): + elif is_base64_encoded_unified_file_id(file_id): managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: raise ProxyException( diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 46f2394c845..ca142088383 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -69,21 +69,25 @@ from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import enforced_model_allowlists from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.route_checks import RouteChecks -from litellm.proxy.auth.user_api_key_auth import ( - _get_bearer_token, +from litellm.proxy.auth.user_api_key_auth import ( # noqa: F401 # legacy module exports + _get_bearer_token, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_bearer_token, is_no_auth_dev_mode, user_api_key_auth, user_api_key_auth_websocket, user_api_key_auth_websocket_for_model, ) from litellm.proxy.common_request_processing import open_sse_before_first_byte -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, - _safe_set_request_parsed_body, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export get_form_data, get_request_body, is_json_content_type, + read_request_body, + safe_get_request_headers, + safe_set_request_parsed_body, ) from litellm.proxy.common_utils.resource_ownership import is_proxy_admin from litellm.proxy.common_utils.sse_keepalive import ( @@ -142,7 +146,7 @@ def create_request_copy(request: Request): return { "method": request.method, "url": str(request.url), - "headers": _safe_get_request_headers(request).copy(), + "headers": safe_get_request_headers(request).copy(), "cookies": request.cookies, "query_params": dict(request.query_params), } @@ -468,7 +472,7 @@ async def fal_ai_proxy_route( status_code=401, detail="FAL_AI_API_KEY is not set and no fal_ai pass-through deployment credentials are configured", ) - if "/requests/" not in endpoint and fal_ai_passthrough_cost(endpoint, await _read_request_body(request)) is None: + if "/requests/" not in endpoint and fal_ai_passthrough_cost(endpoint, await read_request_body(request)) is None: raise HTTPException( status_code=400, detail=f"fal_ai/{endpoint} has no pricing entry for this request; only priced Fal requests can be submitted through /fal_ai", @@ -510,7 +514,7 @@ async def vllm_proxy_route( method=request.method, endpoint=endpoint, request_query_params=request.query_params, - request_headers=_safe_get_request_headers(request), + request_headers=safe_get_request_headers(request), stream=is_streaming_request, content=None, data=None, @@ -665,7 +669,7 @@ async def bespoke_proxy_route( async def _oss_decision_proxy_route( provider: OssDecisionProvider, request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth ) -> Response: - body: Final = TypeAdapter(dict[str, object]).validate_python(await _read_request_body(request)) + body: Final = TypeAdapter(dict[str, object]).validate_python(await read_request_body(request)) try: _ = validate_oss_request(provider, body) except ValueError as exc: @@ -813,7 +817,7 @@ async def milvus_proxy_route( request_body["collectionName"] = vector_store_index # Update the request object with the modified collection name - _safe_set_request_parsed_body(request, request_body) + safe_set_request_parsed_body(request, request_body) vector_store: Final = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry_by_name( vector_store_name=vector_store_name @@ -871,7 +875,7 @@ async def is_streaming_request_fn(request: Request) -> bool: if content_type and "multipart/form-data" in content_type: _request_body = await get_form_data(request) else: - _request_body = await _read_request_body(request) + _request_body = await read_request_body(request) # rebind-ok: pre-existing rebinding on a rename-only line return is_passthrough_request_streaming(_request_body) return False @@ -949,7 +953,7 @@ def is_bedrock_count_tokens_endpoint(endpoint: str) -> bool: return "count_tokens" in endpoint or "count-tokens" in endpoint -def _extract_model_from_bedrock_endpoint(endpoint: str) -> str: +def extract_model_from_bedrock_endpoint(endpoint: str) -> str: """ Extract model name from Bedrock endpoint path. @@ -1029,6 +1033,9 @@ def _extract_model_from_bedrock_endpoint(endpoint: str) -> str: ) from e +_extract_model_from_bedrock_endpoint: Final = extract_model_from_bedrock_endpoint + + async def handle_bedrock_passthrough_router_model( model: str, endpoint: str, @@ -1110,7 +1117,7 @@ async def handle_bedrock_passthrough_router_model( return result except Exception as e: # Use common exception handling - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1228,7 +1235,7 @@ async def bedrock_llm_proxy_route( version, ) - request_body: Final = await _read_request_body(request=request) + request_body: Final = await read_request_body(request=request) if is_bedrock_count_tokens_endpoint(endpoint): return await handle_bedrock_count_tokens( @@ -1241,7 +1248,7 @@ async def bedrock_llm_proxy_route( # Extract model from endpoint path using helper try: - model: Final = _extract_model_from_bedrock_endpoint(endpoint=endpoint) + model: Final = extract_model_from_bedrock_endpoint(endpoint=endpoint) except ValueError as e: raise HTTPException( status_code=400, @@ -1307,7 +1314,7 @@ async def bedrock_llm_proxy_route( return result except Exception as e: - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1645,7 +1652,7 @@ async def azure_speech_proxy_route( target_url: Final = base_url.copy_with( path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, normalized_endpoint_path) ) - request_headers: Final = _safe_get_request_headers(request) + request_headers: Final = safe_get_request_headers(request) upstream_headers: Final = MappingProxyType( { header_name: header_value @@ -1957,8 +1964,10 @@ async def assemblyai_proxy_route( [Docs](https://api.assemblyai.com) """ # Set base URL based on the route - assembly_region: Final = AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(url=str(request.url)) - base_target_url = AssemblyAIPassthroughLoggingHandler._get_assembly_base_url_from_region(region=assembly_region) + assembly_region: Final = AssemblyAIPassthroughLoggingHandler.get_assembly_region_from_url(url=str(request.url)) + base_target_url: Final = AssemblyAIPassthroughLoggingHandler.get_assembly_base_url_from_region( + region=assembly_region + ) encoded_endpoint = httpx.URL(endpoint).path # Ensure endpoint starts with '/' for proper URL construction if not encoded_endpoint.startswith("/"): @@ -2092,7 +2101,7 @@ async def _relay_router_model( method=request.method, endpoint=endpoint, request_query_params=request.query_params, - request_headers=_safe_get_request_headers(request), + request_headers=safe_get_request_headers(request), stream=is_streaming_request, content=None, data=None, @@ -2332,7 +2341,7 @@ async def azure_proxy_route( base_target_url = _optional_str(litellm_params.get("api_base")) if base_target_url is None: raise Exception(f"API base not found for {part}") - return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler( + return await BaseOpenAIPassThroughHandler.base_openai_pass_through_handler( endpoint=endpoint, request=request, fastapi_response=fastapi_response, @@ -2360,7 +2369,7 @@ async def azure_proxy_route( if azure_api_key is None: raise Exception("Required 'AZURE_API_KEY' in environment to make pass-through calls to Azure.") - return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler( + return await BaseOpenAIPassThroughHandler.base_openai_pass_through_handler( endpoint=endpoint, request=request, fastapi_response=fastapi_response, @@ -2431,7 +2440,7 @@ def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict: Returns: dict: Headers dictionary with only allowed headers """ - incoming_headers: Final = _safe_get_request_headers(request) + incoming_headers: Final = safe_get_request_headers(request) headers: Final = {} for header_name in ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: if header_name in incoming_headers: @@ -2531,7 +2540,7 @@ def _normalize_credential_value(value: str) -> str: with no recognized scheme prefix, so a bare token (or a real Google credential that carries no scheme) falls back to its own value. """ - return _get_bearer_token(value) or value + return get_bearer_token(value) or value _VERTEX_UPSTREAM_CREDENTIAL_HEADERS: Final = frozenset({"authorization", "x-goog-api-key"}) @@ -2616,7 +2625,7 @@ def _is_authenticated_caller_secret(value: str, user_api_key_dict: UserAPIKeyAut def _caller_headers_without_litellm_secrets( request: Request, user_api_key_dict: UserAPIKeyAuth, never_forwarded: frozenset[str] ) -> Mapping[str, str]: - incoming: Final = _safe_get_request_headers(request) + incoming: Final = safe_get_request_headers(request) dropped_by_name: Final = never_forwarded.union( (_MAPPED_ROUTE_CALLER_KEY_HEADER, *_operator_configured_caller_key_header_names()) ) @@ -3048,7 +3057,7 @@ async def openai_proxy_route( if openai_api_key is None: raise Exception("Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.") - return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler( + return await BaseOpenAIPassThroughHandler.base_openai_pass_through_handler( endpoint=endpoint, request=request, fastapi_response=fastapi_response, @@ -3320,7 +3329,7 @@ async def deepgram_listen_websocket_route( class BaseOpenAIPassThroughHandler: @staticmethod - async def _base_openai_pass_through_handler( + async def base_openai_pass_through_handler( endpoint: str, request: Request, fastapi_response: Response, @@ -3372,12 +3381,14 @@ class BaseOpenAIPassThroughHandler: return received_value + _base_openai_pass_through_handler = base_openai_pass_through_handler + @staticmethod def _append_openai_beta_header(headers: dict, request: Request) -> dict: """ Appends the OpenAI-Beta header to the headers if the request is an OpenAI Assistants API request """ - if RouteChecks._is_assistants_api_request(request) is True and "OpenAI-Beta" not in headers: + if RouteChecks.is_assistants_api_request(request) is True and "OpenAI-Beta" not in headers: headers["OpenAI-Beta"] = "assistants=v2" return headers @@ -4023,7 +4034,7 @@ async def handle_gigachat_passthrough_router_model( is_streaming: Final = request_body.get("stream", False) # pyright: ignore[reportUnknownVariableType] # request_body is dict[Unknown, Unknown] - data: Final[dict[str, object]] = await _read_request_body(request=request) + data: Final[dict[str, object]] = await read_request_body(request=request) if user_api_key_dict is not None: auth_metadata: Final = { metadata_key: value @@ -4096,7 +4107,7 @@ async def handle_gigachat_passthrough_router_model( ) except Exception as e: # noqa: BLE001 # Safe catch-all for handle exception # Use common exception handling - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index dfb731972ac..73d40369d30 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -173,7 +173,7 @@ class AnthropicPassthroughLoggingHandler: all_chunks: Sequence[str | bytes], model: str, speed: str | None ) -> ModelResponse | None: try: - return AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks( + return AnthropicPassthroughLoggingHandler.build_usage_only_response_from_chunks( all_chunks=all_chunks, model=model, speed=speed ) except Exception as e: # noqa: BLE001 # the usage-only fallback must never raise out of failure logging @@ -188,7 +188,7 @@ class AnthropicPassthroughLoggingHandler: speed: str | None, ) -> ModelResponse | TextCompletionResponse | None: try: - assembled: Final = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( + assembled: Final = AnthropicPassthroughLoggingHandler.build_complete_streaming_response( all_chunks=all_chunks, litellm_logging_obj=litellm_logging_obj, model=model, @@ -467,7 +467,7 @@ class AnthropicPassthroughLoggingHandler: return kwargs @staticmethod - def _handle_logging_anthropic_collected_chunks( + def handle_logging_anthropic_collected_chunks( litellm_logging_obj: LiteLLMLoggingObj, passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, @@ -514,6 +514,8 @@ class AnthropicPassthroughLoggingHandler: "kwargs": kwargs, } + _handle_logging_anthropic_collected_chunks = handle_logging_anthropic_collected_chunks + @staticmethod def _split_sse_chunk_into_events(chunk: str | bytes) -> list[str]: """ @@ -539,7 +541,7 @@ class AnthropicPassthroughLoggingHandler: return events @staticmethod - def _build_complete_streaming_response( + def build_complete_streaming_response( all_chunks: Sequence[str | bytes], litellm_logging_obj: LiteLLMLoggingObj, model: str, @@ -577,6 +579,8 @@ class AnthropicPassthroughLoggingHandler: speed=speed, ) + _build_complete_streaming_response = build_complete_streaming_response + # Anthropic SSE block/delta types that the fast path is NOT allowed to # collapse -- their presence forces the unchanged legacy path so tool # calls, thinking, citations, etc. keep byte-identical reconstruction. @@ -775,7 +779,7 @@ class AnthropicPassthroughLoggingHandler: return None @staticmethod - def _build_usage_only_response_from_chunks( + def build_usage_only_response_from_chunks( all_chunks: Sequence[str | bytes], model: str, speed: str | None = None, @@ -894,6 +898,8 @@ class AnthropicPassthroughLoggingHandler: usage=usage_obj, ) + _build_usage_only_response_from_chunks = build_usage_only_response_from_chunks + @staticmethod def batch_creation_handler( httpx_response: httpx.Response, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py index 812f72faecc..c9c89870610 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py @@ -147,7 +147,7 @@ class AssemblyAIPassthroughLoggingHandler: logging_obj.model_call_details["response_cost"] = response_cost asyncio.run( - pass_through_endpoint_logging._handle_logging( + pass_through_endpoint_logging.handle_logging( logging_obj=logging_obj, standard_logging_response_object=self._get_response_to_log(transcript_response), result=result, @@ -216,7 +216,7 @@ class AssemblyAIPassthroughLoggingHandler: """ for _ in range(self.max_polling_attempts): # 180 attempts * 10s = 30 minutes max transcript = self._get_assembly_transcript( - request_region=AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(url=url_route), + request_region=AssemblyAIPassthroughLoggingHandler.get_assembly_region_from_url(url=url_route), transcript_id=transcript_id, ) if transcript is None: @@ -279,14 +279,16 @@ class AssemblyAIPassthroughLoggingHandler: return None @staticmethod - def _should_log_request(request_method: str) -> bool: + def should_log_request(request_method: str) -> bool: """ only POST transcription jobs are logged. litellm will POLL assembly to wait for the transcription to complete to log the complete response / cost """ return request_method == "POST" + _should_log_request = should_log_request + @staticmethod - def _get_assembly_region_from_url(url: str | None) -> Literal["eu"] | None: + def get_assembly_region_from_url(url: str | None) -> Literal["eu"] | None: """ Get the region from the URL """ @@ -296,8 +298,10 @@ class AssemblyAIPassthroughLoggingHandler: return "eu" return None + _get_assembly_region_from_url = get_assembly_region_from_url + @staticmethod - def _get_assembly_base_url_from_region(region: Literal["eu"] | None) -> str: + def get_assembly_base_url_from_region(region: Literal["eu"] | None) -> str: """ Get the base URL for the AssemblyAI API if region == "eu", return "https://api.eu.assemblyai.com" @@ -306,3 +310,5 @@ class AssemblyAIPassthroughLoggingHandler: if region == "eu": return "https://api.eu.assemblyai.com" return "https://api.assemblyai.com" + + _get_assembly_base_url_from_region = get_assembly_base_url_from_region diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py index 0fcc1ccda38..2e9d92ed309 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py @@ -78,7 +78,7 @@ def _is_openai_compatible_host(hostname: str | None) -> bool: return _hostname_matches(hostname, _OPENAI_HOSTNAMES) or _hostname_matches(hostname, _AZURE_OPENAI_HOSTNAMES) -def _is_openai_compatible_url(url_route: str | None) -> bool: +def is_openai_compatible_url(url_route: str | None) -> bool: """True if the URL targets an OpenAI-compatible API surface. For the shared Azure Cognitive Services domains we additionally require an @@ -99,6 +99,9 @@ def _is_openai_compatible_url(url_route: str | None) -> bool: return False +_is_openai_compatible_url: Final = is_openai_compatible_url + + def _is_remote_high_detail_image(part: object) -> bool: if not isinstance(part, Mapping) or part.get("type") != "image_url": return False @@ -608,7 +611,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): return None @staticmethod - def _handle_logging_openai_collected_chunks( + def handle_logging_openai_collected_chunks( litellm_logging_obj: LiteLLMLoggingObj, passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, @@ -725,3 +728,5 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): "result": None, "kwargs": {}, } + + _handle_logging_openai_collected_chunks = handle_logging_openai_collected_chunks diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 47e1f21d1a9..cfa569b5219 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -561,7 +561,7 @@ class VertexPassthroughLoggingHandler: } @staticmethod - def _handle_logging_vertex_collected_chunks( + def handle_logging_vertex_collected_chunks( litellm_logging_obj: LiteLLMLoggingObj, passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, @@ -616,6 +616,8 @@ class VertexPassthroughLoggingHandler: "kwargs": kwargs, } + _handle_logging_vertex_collected_chunks = handle_logging_vertex_collected_chunks + @staticmethod def _build_complete_streaming_response( all_chunks: list[str], diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index fc70ffe8697..07053ad0f32 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -91,9 +91,11 @@ from litellm.proxy.common_request_processing import ( resolve_litellm_call_id, ) from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_body_call_id, with_call_id -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, + safe_get_request_headers, ) from litellm.proxy.common_utils.openai_error_payload import ( LITELLM_CALL_ID_HEADER, @@ -105,11 +107,12 @@ from litellm.proxy.common_utils.openai_error_payload import ( from litellm.proxy.common_utils.sse_keepalive import ( wrap_passthrough_sse_bytes_with_keepalive_pings, ) -from litellm.proxy.litellm_pre_call_utils import ( +from litellm.proxy.litellm_pre_call_utils import ( # noqa: F401 # legacy module exports LiteLLMProxyRequestSetup, - _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above + _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export _key_or_team_allows_client_pricing_override, # pyright: ignore[reportPrivateUsage] # reuse the proxy's pricing trust policy _strip_client_pricing_overrides, # pyright: ignore[reportPrivateUsage] # sanitize before trusted hooks add guardrail costs + get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path @@ -586,7 +589,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): ) @staticmethod - def _init_kwargs_for_pass_through_endpoint( + def init_kwargs_for_pass_through_endpoint( request: Request, user_api_key_dict: UserAPIKeyAuth, passthrough_logging_payload: PassthroughStandardLoggingPayload, @@ -697,6 +700,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): return kwargs + _init_kwargs_for_pass_through_endpoint = init_kwargs_for_pass_through_endpoint + @staticmethod def construct_target_url_with_subpath(base_target: str, subpath: str, include_subpath: bool | None) -> str: """ @@ -768,7 +773,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): return combined @staticmethod - def _update_stream_param_based_on_request_body( + def update_stream_param_based_on_request_body( parsed_body: dict, stream: bool | None = None, ) -> bool | None: @@ -780,6 +785,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): return parsed_body.get("stream", stream) return stream + _update_stream_param_based_on_request_body = update_stream_param_based_on_request_body + def _carry_guardrail_logging_info(request_data: dict, guardrail_data: dict | None) -> None: """Copy guardrail logging entries from ``guardrail_data`` onto ``request_data``. @@ -863,7 +870,7 @@ def _resolve_team_callback_wiring( otherwise reject the vars mid-request. """ try: - callback_settings_obj: Final = _get_dynamic_logging_metadata( + callback_settings_obj: Final = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) if callback_settings_obj and callback_settings_obj.callback_vars: @@ -1125,7 +1132,7 @@ async def pass_through_request( url = httpx.URL(target) headers = custom_headers headers = HttpPassThroughEndpointHelpers.forward_headers_from_request( - request_headers=_safe_get_request_headers(request).copy(), + request_headers=safe_get_request_headers(request).copy(), headers=headers, forward_headers=forward_headers, ) @@ -1161,7 +1168,7 @@ async def pass_through_request( # Don't parse multipart body here - it will be handled by make_multipart_http_request _parsed_body = {} else: - _parsed_body = await _read_request_body(request) + _parsed_body = await read_request_body(request) # rebind-ok: pre-existing rebinding on a rename-only line verbose_proxy_logger.debug( "Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n", url, @@ -1258,7 +1265,7 @@ async def pass_through_request( request_method=getattr(request, "method", None), cost_per_request=cost_per_request, ) - kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + kwargs = HttpPassThroughEndpointHelpers.init_kwargs_for_pass_through_endpoint( # rebind-ok: pre-existing rebinding on a rename-only line user_api_key_dict=user_api_key_dict, _parsed_body=_parsed_body, passthrough_logging_payload=passthrough_logging_payload, @@ -1432,7 +1439,7 @@ async def pass_through_request( "headers": upstream_headers, }, ) - stream = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body( + stream = HttpPassThroughEndpointHelpers.update_stream_param_based_on_request_body( parsed_body=_parsed_body or {}, stream=stream, ) @@ -1980,7 +1987,7 @@ def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> di # Only add tags key if there are tags to add if tags_to_add: - metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags( + metadata["tags"] = LiteLLMProxyRequestSetup.merge_tags( request_tags=metadata.get("tags"), tags_to_add=tags_to_add, ) @@ -2480,7 +2487,7 @@ async def websocket_passthrough_request( ) # Initialize kwargs for logging using the same pattern as HTTP passthrough - kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + kwargs: Final = HttpPassThroughEndpointHelpers.init_kwargs_for_pass_through_endpoint( user_api_key_dict=user_api_key_dict, _parsed_body={}, # WebSocket doesn't have a traditional request body passthrough_logging_payload=passthrough_logging_payload, diff --git a/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py b/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py index de9b0acf081..1c5d270542f 100644 --- a/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py +++ b/litellm/proxy/pass_through_endpoints/passthrough_guardrails.py @@ -228,7 +228,7 @@ class PassthroughGuardrailHandler: Dict of guardrail names to run (format: {guardrail_name: True}), or None """ from litellm.proxy.litellm_pre_call_utils import ( - _add_guardrails_from_key_or_team_metadata, + add_guardrails_from_key_or_team_metadata, ) # Normalize config to dict format (handles both list and dict) @@ -252,7 +252,7 @@ class PassthroughGuardrailHandler: # Add org/team/key level guardrails using shared helper temp_data: Final[dict[str, Any]] = {"metadata": {}} - _add_guardrails_from_key_or_team_metadata( + add_guardrails_from_key_or_team_metadata( key_metadata=user_api_key_dict.metadata, team_metadata=user_api_key_dict.team_metadata, data=temp_data, diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index 8199a96ceab..8f6a688fecd 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -147,7 +147,7 @@ class PassThroughStreamingHandler: route_streaming_logging: RouteStreamingLogging | None = None, ): resolved_route_streaming_logging: Final[RouteStreamingLogging] = ( - route_streaming_logging or PassThroughStreamingHandler._route_streaming_logging_to_handler + route_streaming_logging or PassThroughStreamingHandler.route_streaming_logging_to_handler ) raw_bytes: Final[list[bytes]] = [] resolved_request_body: Final[dict[str, object]] = request_body or {} @@ -219,7 +219,7 @@ class PassThroughStreamingHandler: PassThroughStreamingHandler._stamp_first_chunk_if_needed(litellm_logging_obj) complete_frames, pending = split_complete_sse_frames(pending + chunk) if complete_frames: - yield ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection( + yield ProxyBaseLLMRequestProcessing.process_chunk_with_cost_injection( complete_frames, resolved_model_name, litellm_logging_obj ) if pending: @@ -275,7 +275,7 @@ class PassThroughStreamingHandler: bind_budget_reservation_to_callbacks(litellm_logging_obj.litellm_params) @staticmethod - async def _route_streaming_logging_to_handler( + async def route_streaming_logging_to_handler( litellm_logging_obj: LiteLLMLoggingObj, passthrough_success_handler_obj: PassThroughEndpointLogging, url_route: str, @@ -373,6 +373,8 @@ class PassThroughStreamingHandler: except Exception as e: verbose_proxy_logger.error("Error in _route_streaming_logging_to_handler: %s", e) + _route_streaming_logging_to_handler = route_streaming_logging_to_handler + @staticmethod def _build_passthrough_logging_result( litellm_logging_obj: LiteLLMLoggingObj, @@ -397,7 +399,7 @@ class PassThroughStreamingHandler: kwargs: dict = {} if endpoint_type == EndpointType.ANTHROPIC: anthropic_passthrough_logging_handler_result: Final = ( - AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks( + AnthropicPassthroughLoggingHandler.handle_logging_anthropic_collected_chunks( litellm_logging_obj=litellm_logging_obj, passthrough_success_handler_obj=passthrough_success_handler_obj, url_route=url_route, @@ -412,7 +414,7 @@ class PassThroughStreamingHandler: kwargs = anthropic_passthrough_logging_handler_result["kwargs"] elif endpoint_type == EndpointType.VERTEX_AI: vertex_passthrough_logging_handler_result: Final = ( - VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks( + VertexPassthroughLoggingHandler.handle_logging_vertex_collected_chunks( litellm_logging_obj=litellm_logging_obj, passthrough_success_handler_obj=passthrough_success_handler_obj, url_route=url_route, @@ -448,7 +450,7 @@ class PassThroughStreamingHandler: ) elif endpoint_type == EndpointType.OPENAI: openai_passthrough_logging_handler_result: Final = ( - OpenAIPassthroughLoggingHandler._handle_logging_openai_collected_chunks( + OpenAIPassthroughLoggingHandler.handle_logging_openai_collected_chunks( litellm_logging_obj=litellm_logging_obj, passthrough_success_handler_obj=passthrough_success_handler_obj, url_route=url_route, diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 65a199d38b6..9c6c3111209 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -117,9 +117,9 @@ class PassThroughEndpointLogging: @property def _log_dispatch(self) -> PassThroughLogDispatch: - return self._injected_log_dispatch if self._injected_log_dispatch is not None else self._handle_logging + return self._injected_log_dispatch if self._injected_log_dispatch is not None else self.handle_logging - async def _handle_logging( + async def handle_logging( self, logging_obj: LiteLLMLoggingObj, standard_logging_response_object: StandardPassThroughResponseObject @@ -130,7 +130,7 @@ class PassThroughEndpointLogging: end_time: datetime, cache_hit: bool, **kwargs, - ): + ) -> None: """Log pass-through success via the shared async dispatch path.""" # Always reached from pass_through_async_success_handler, which runs in # an async context. call_type is "pass_through_endpoint" here, so the @@ -148,6 +148,8 @@ class PassThroughEndpointLogging: **kwargs, ) + _handle_logging = handle_logging + def normalize_llm_passthrough_logging_payload( self, httpx_response: httpx.Response, @@ -445,7 +447,7 @@ class PassThroughEndpointLogging: ) return if self.is_assemblyai_route(url_route) and not self.is_azure_speech_route(custom_llm_provider): - if AssemblyAIPassthroughLoggingHandler._should_log_request(httpx_response.request.method) is not True: + if AssemblyAIPassthroughLoggingHandler.should_log_request(httpx_response.request.method) is not True: return self.assemblyai_passthrough_logging_handler.assemblyai_passthrough_logging_handler( httpx_response=httpx_response, @@ -606,10 +608,10 @@ class PassThroughEndpointLogging: if not url_route: return False from .llm_provider_handlers.openai_passthrough_logging_handler import ( - _is_openai_compatible_url, + is_openai_compatible_url, ) - return _is_openai_compatible_url(url_route) + return is_openai_compatible_url(url_route) def is_gemini_route(self, url_route: str, custom_llm_provider: str | None = None): """Check if the URL route is a Gemini API route.""" diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index efe7a5002e0..cfc1f481988 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -195,7 +195,7 @@ class PolicyRegistry: for policy_name, policy_data in policies_config.items(): try: - policy = self._parse_policy(policy_name, policy_data) + policy = self.parse_policy(policy_name, policy_data) self._policies[policy_name] = policy verbose_proxy_logger.debug("Loaded policy: %s", policy_name) except Exception as e: @@ -207,7 +207,7 @@ class PolicyRegistry: self._initialized = True verbose_proxy_logger.info("Loaded %s policies", len(self._policies)) - def _parse_policy(self, policy_name: str, policy_data: dict[str, Any]) -> Policy: + def parse_policy(self, policy_name: str, policy_data: dict[str, Any]) -> Policy: """ Parse a policy from raw configuration data. @@ -246,6 +246,8 @@ class PolicyRegistry: pipeline=pipeline, ) + _parse_policy = parse_policy + @staticmethod def _parse_pipeline( pipeline_data: Optional["_RawPipelineConfig"], @@ -427,7 +429,7 @@ class PolicyRegistry: created_policy: Final = await _policy_table(prisma_client).create(data=data) # Also add to in-memory registry - policy: Final = self._parse_policy( + policy: Final = self.parse_policy( policy_request.policy_name, { "inherit": policy_request.inherit, @@ -648,7 +650,7 @@ class PolicyRegistry: try: production: Final = await self.get_all_policies_from_db(prisma_client, version_status="production") db_policies: Final = { - policy_response.policy_name: self._parse_policy( + policy_response.policy_name: self.parse_policy( policy_response.policy_name, { "inherit": policy_response.inherit, @@ -679,7 +681,7 @@ class PolicyRegistry: order={"created_at": "desc"}, ) for row in non_production: - policy = self._parse_policy( + policy = self.parse_policy( row.policy_name, { "inherit": row.inherit, @@ -731,7 +733,7 @@ class PolicyRegistry: # Build a temporary in-memory map for resolution temp_policies: Final = {} for policy_response in policies: - policy = self._parse_policy( + policy = self.parse_policy( policy_response.policy_name, { "inherit": policy_response.inherit, @@ -959,7 +961,7 @@ class PolicyRegistry: # Update in-memory registry: remove old production (by name), add this one self.remove_policy(policy_name) - policy: Final = self._parse_policy( + policy: Final = self.parse_policy( policy_name, { "inherit": updated.inherit, diff --git a/litellm/proxy/policy_engine/policy_validator.py b/litellm/proxy/policy_engine/policy_validator.py index 17542ba814d..36263ebd26e 100644 --- a/litellm/proxy/policy_engine/policy_validator.py +++ b/litellm/proxy/policy_engine/policy_validator.py @@ -193,7 +193,7 @@ class PolicyValidator: # A concrete entry is one the request-time matcher compares by exact equality; # only a trailing "*" is a wildcard (RouteChecks._is_wildcard_pattern), and those # are left unvalidated since they may match zero entities today and more later. - is_pattern: Final = RouteChecks._is_wildcard_pattern + is_pattern: Final = RouteChecks.is_wildcard_pattern concrete_teams: Final = [t for t in (teams or []) if not is_pattern(pattern=t)] concrete_keys: Final = [k for k in (keys or []) if not is_pattern(pattern=k)] concrete_models: Final = [m for m in (models or []) if not is_pattern(pattern=m)] @@ -429,7 +429,7 @@ class PolicyValidator: for policy_name, policy_data in policy_config.items(): try: - policy = temp_registry._parse_policy(policy_name, policy_data) + policy = temp_registry.parse_policy(policy_name, policy_data) policies[policy_name] = policy except Exception as e: errors.append( diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 53a447bc6ec..c7371f294f0 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -204,18 +204,22 @@ def deprecated_v2_flag_passed_on_cli() -> bool: class ProxyInitializationHelpers: @staticmethod - def _echo_litellm_version(): + def echo_litellm_version(): pkg_version: Final = importlib.metadata.version("litellm") click.echo(f"\nLiteLLM: Current Version = {pkg_version}\n") + _echo_litellm_version = echo_litellm_version + @staticmethod - def _run_health_check(host, port): + def run_health_check(host, port): print("\nLiteLLM: Health Testing models in config") response: Final = httpx.get(url=f"http://{host}:{port}/health") print(json.dumps(response.json(), indent=4)) + _run_health_check = run_health_check + @staticmethod - def _run_config_validation(config: str | None) -> None: + def run_config_validation(config: str | None) -> None: if config is None: raise click.UsageError("--validate_config requires --config ") import asyncio @@ -233,8 +237,10 @@ class ProxyInitializationHelpers: raise click.exceptions.Exit(1) from error click.echo(f"LiteLLM: config OK ({model_count} models)") + _run_config_validation = run_config_validation + @staticmethod - def _run_test_chat_completion( + def run_test_chat_completion( host: str, port: int, model: str, @@ -283,8 +289,10 @@ class ProxyInitializationHelpers: ) print(completion_response) + _run_test_chat_completion = run_test_chat_completion + @staticmethod - def _get_default_unvicorn_init_args( + def get_default_unvicorn_init_args( host: str, port: int, log_config: str | None = None, @@ -328,8 +336,10 @@ class ProxyInitializationHelpers: ) return uvicorn_args + _get_default_unvicorn_init_args = get_default_unvicorn_init_args + @staticmethod - def _apply_uvicorn_max_requests_jitter( + def apply_uvicorn_max_requests_jitter( uvicorn_args: dict, max_requests_before_restart: int | None, jitter: int, @@ -356,6 +366,8 @@ class ProxyInitializationHelpers: f"Ignoring the flag.\033[0m" ) + _apply_uvicorn_max_requests_jitter = apply_uvicorn_max_requests_jitter + @staticmethod def _get_reload_options(config_path: str | None) -> dict: """Build uvicorn reload kwargs so --reload also reacts to .env and YAML edits.""" @@ -419,7 +431,7 @@ class ProxyInitializationHelpers: return True @staticmethod - def _configure_dev_reload(uvicorn_args: dict, config_path: str | None) -> None: + def configure_dev_reload(uvicorn_args: dict, config_path: str | None) -> None: """Wire up --reload (dev only): watch *.py, the --config YAML, and .env, and signal reloaded workers to re-read .env with override so edits to existing keys actually take effect rather than staying masked by the @@ -436,8 +448,10 @@ class ProxyInitializationHelpers: "to let a shell-exported value take precedence." ) + _configure_dev_reload = configure_dev_reload + @staticmethod - def _init_hypercorn_server( + def init_hypercorn_server( app: FastAPI, host: str, port: int, @@ -469,8 +483,10 @@ class ProxyInitializationHelpers: # hypercorn serve raises a type warning when passing a fast api app - even though fast API is a valid type asyncio.run(serve(app, config)) + _init_hypercorn_server = init_hypercorn_server + @staticmethod - def _init_granian_server( + def init_granian_server( host: str, port: int, num_workers: int, @@ -519,8 +535,10 @@ class ProxyInitializationHelpers: Granian(**kwargs).serve() + _init_granian_server = init_granian_server + @staticmethod - def _run_gunicorn_server( + def run_gunicorn_server( host: str, port: int, app: FastAPI, @@ -635,8 +653,10 @@ class ProxyInitializationHelpers: start_query_engine_reaper() StandaloneApplication(app=app, options=gunicorn_options).run() # Run gunicorn + _run_gunicorn_server = run_gunicorn_server + @staticmethod - def _run_ollama_serve(): + def run_ollama_serve(): try: command: Final = ["ollama", "serve"] @@ -647,20 +667,26 @@ class ProxyInitializationHelpers: LiteLLM Warning: proxy started with `ollama` model\n`ollama serve` failed with Exception{e}. \nEnsure you run `ollama serve` """) + _run_ollama_serve = run_ollama_serve + @staticmethod - def _is_port_in_use(port): + def is_port_in_use(port): import socket with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: return s.connect_ex(("localhost", port)) == 0 + _is_port_in_use = is_port_in_use + @staticmethod - def _get_loop_type(): + def get_loop_type(): """Helper function to determine the event loop type based on platform""" if sys.platform in ("win32", "cygwin", "cli"): return None # Let uvicorn choose the default loop on Windows return "uvloop" + _get_loop_type = get_loop_type + @staticmethod def _prometheus_callback_configured(litellm_settings: Mapping[str, object] | None) -> bool: if litellm_settings is None: @@ -676,7 +702,7 @@ class ProxyInitializationHelpers: ) @staticmethod - def _maybe_setup_prometheus_multiproc_dir( + def maybe_setup_prometheus_multiproc_dir( num_workers: int, litellm_settings: dict | None, prometheus_metrics_port: int | None = None, @@ -707,6 +733,8 @@ class ProxyInitializationHelpers: print(f"LiteLLM: {action} PROMETHEUS_MULTIPROC_DIR={multiproc_dir}") return multiproc_dir + _maybe_setup_prometheus_multiproc_dir = maybe_setup_prometheus_multiproc_dir + @click.command() @click.argument("cli_args", nargs=-1) @@ -1102,10 +1130,10 @@ def run_server( except ModuleNotFoundError as e: raise ModuleNotFoundError(f"Missing dependency {e}. Run `pip install 'litellm[proxy]'`") from e if version is True: - ProxyInitializationHelpers._echo_litellm_version() + ProxyInitializationHelpers.echo_litellm_version() return if validate_config is True: - ProxyInitializationHelpers._run_config_validation(config) + ProxyInitializationHelpers.run_config_validation(config) return if enforce_prisma_migration_check: print( @@ -1114,12 +1142,12 @@ def run_server( "when database setup fails at startup. You can safely remove it.\033[0m" ) if model and "ollama" in model and api_base is None: - ProxyInitializationHelpers._run_ollama_serve() + ProxyInitializationHelpers.run_ollama_serve() if health is True: - ProxyInitializationHelpers._run_health_check(host, port) + ProxyInitializationHelpers.run_health_check(host, port) return if test is True: - ProxyInitializationHelpers._run_test_chat_completion(host, port, model, test) + ProxyInitializationHelpers.run_test_chat_completion(host, port, model, test) return else: if headers: @@ -1491,7 +1519,7 @@ def run_server( ) sys.exit(1) export_pooled_database_url(pooled_database_url) - if port == 4000 and ProxyInitializationHelpers._is_port_in_use(port): + if port == 4000 and ProxyInitializationHelpers.is_port_in_use(port): port = random.randint(1024, 49152) if prometheus_metrics_port == port: raise click.UsageError("--prometheus_metrics_port must differ from --port") @@ -1512,7 +1540,7 @@ def run_server( return # Auto-create PROMETHEUS_MULTIPROC_DIR for multi-worker setups - prometheus_multiproc_dir: Final = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir( + prometheus_multiproc_dir: Final = ProxyInitializationHelpers.maybe_setup_prometheus_multiproc_dir( num_workers=num_workers, litellm_settings=litellm_settings if config else None, prometheus_metrics_port=prometheus_metrics_port, @@ -1533,7 +1561,7 @@ def run_server( ) running_uvicorn: Final = run_gunicorn is False and run_hypercorn is False - uvicorn_args: Final = ProxyInitializationHelpers._get_default_unvicorn_init_args( + uvicorn_args: Final = ProxyInitializationHelpers.get_default_unvicorn_init_args( host=host, port=port, log_config=log_config, @@ -1547,7 +1575,7 @@ def run_server( if limit_concurrency is not None: uvicorn_args["limit_concurrency"] = limit_concurrency if max_requests_before_restart_jitter is not None: - ProxyInitializationHelpers._apply_uvicorn_max_requests_jitter( + ProxyInitializationHelpers.apply_uvicorn_max_requests_jitter( uvicorn_args=uvicorn_args, max_requests_before_restart=max_requests_before_restart, jitter=max_requests_before_restart_jitter, @@ -1559,12 +1587,12 @@ def run_server( uvicorn_args["ssl_keyfile"] = ssl_keyfile_path uvicorn_args["ssl_certfile"] = ssl_certfile_path - loop_type: Final = ProxyInitializationHelpers._get_loop_type() + loop_type: Final = ProxyInitializationHelpers.get_loop_type() if loop_type: uvicorn_args["loop"] = loop_type if reload: - ProxyInitializationHelpers._configure_dev_reload(uvicorn_args, config) + ProxyInitializationHelpers.configure_dev_reload(uvicorn_args, config) if num_workers > 1: start_query_engine_reaper() @@ -1573,7 +1601,7 @@ def run_server( workers=num_workers, ) elif run_gunicorn is True: - ProxyInitializationHelpers._run_gunicorn_server( + ProxyInitializationHelpers.run_gunicorn_server( host=host, port=port, app=app, @@ -1584,7 +1612,7 @@ def run_server( max_requests_before_restart_jitter=max_requests_before_restart_jitter, ) elif run_hypercorn is True: - ProxyInitializationHelpers._init_hypercorn_server( + ProxyInitializationHelpers.init_hypercorn_server( app=app, host=host, port=port, @@ -1593,7 +1621,7 @@ def run_server( ciphers=ciphers, ) elif run_granian is True: - ProxyInitializationHelpers._init_granian_server( + ProxyInitializationHelpers.init_granian_server( host=host, port=port, num_workers=num_workers, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index aa00bfe5d26..8afacda2e4c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -133,7 +133,10 @@ from litellm.proxy.common_utils.callback_utils import ( process_callback, strip_callback_config, ) -from litellm.proxy.common_utils.realtime_utils import _realtime_request_body +from litellm.proxy.common_utils.realtime_utils import ( # noqa: F401, RUF100 # legacy module exports + _realtime_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + realtime_request_body, +) from litellm.proxy.management_helpers.auto_router_availability import AutoRouterCatalogEntry, build_auto_router_catalog from litellm.router_utils.access_windows import access_windows_config_error from litellm.router_utils.add_retry_fallback_headers import ( @@ -384,8 +387,9 @@ from litellm.proxy.auth.model_checks import ( get_team_models, ) from litellm.proxy.auth.password_policy import validate_password_not_breached, validate_password_policy -from litellm.proxy.auth.user_api_key_auth import ( - _fetch_global_spend_with_event_coordination, +from litellm.proxy.auth.user_api_key_auth import ( # noqa: F401, RUF100 # legacy module exports + _fetch_global_spend_with_event_coordination, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + fetch_global_spend_with_event_coordination, user_api_key_auth, user_api_key_auth_websocket, ) @@ -394,17 +398,19 @@ from litellm.proxy.bug_report_config import build_proxy_bug_report ## Import All Misc routes here ## from litellm.proxy.caching_routes import router as caching_router -from litellm.proxy.common_request_processing import ( +from litellm.proxy.common_request_processing import ( # noqa: F401, RUF100 # legacy module exports KNOWN_PROXY_ROUTES, ProxyBaseLLMRequestProcessing, - _is_azure_model_router_request, - _should_return_raw_model_name, + _is_azure_model_router_request, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _should_return_raw_model_name, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export close_guarded_stream, create_response, + is_azure_model_router_request, log_llm_api_exception, open_sse_before_first_byte, request_litellm_call_id, resolve_litellm_call_id, + should_return_raw_model_name, ttft_keepalive_interval, ) from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( @@ -437,12 +443,14 @@ from litellm.proxy.common_utils.healthy_model_filter import ( ) from litellm.proxy.common_utils.html_forms.default_credentials_hint import should_hide_default_credentials_hint from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401, RUF100 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export check_file_size_under_limit, get_form_data, + read_request_body, resolve_inference_model, + safe_get_request_headers, ) from litellm.proxy.common_utils.load_config_utils import get_config_from_bucket from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations @@ -570,14 +578,20 @@ from litellm.proxy.health_check import ( perform_health_check, ) from litellm.proxy.health_endpoints._health_endpoints import router as health_router -from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, +from litellm.proxy.hooks.model_max_budget_limiter import ( # noqa: F401, RUF100 # legacy module exports + PROXY_VirtualKeyModelMaxBudgetLimiter, + _PROXY_VirtualKeyModelMaxBudgetLimiter, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) from litellm.proxy.hooks.parallel_request_limiter_v3 import fail_closed_rate_limit_enforcement_enabled -from litellm.proxy.hooks.prompt_injection_detection import ( - _OPTIONAL_PromptInjectionDetection, +from litellm.proxy.hooks.prompt_injection_detection import ( # noqa: F401, RUF100 # legacy module exports + OPTIONAL_PromptInjectionDetection, + _OPTIONAL_PromptInjectionDetection, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) +from litellm.proxy.hooks.proxy_track_cost_callback import ( # noqa: F401, RUF100 # legacy module exports + ProxyDBLogger, + _ProxyDBLogger, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + run_spend_event, ) -from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_spend_event from litellm.proxy.image_endpoints.endpoints import router as image_router from litellm.proxy.lens.dataset_endpoints import router as lens_dataset_router from litellm.proxy.lens.endpoints import router as lens_router @@ -612,10 +626,12 @@ from litellm.proxy.management_endpoints.cache_settings_endpoints import ( from litellm.proxy.management_endpoints.callback_management_endpoints import ( router as callback_management_endpoints_router, ) -from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_privileges, - _user_has_admin_view, +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401, RUF100 # legacy module exports + _user_has_admin_privileges, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export admin_can_invite_user, + user_api_key_has_admin_view, + user_has_admin_privileges, ) from litellm.proxy.management_endpoints.coordination_redis_endpoints import ( get_persisted_coordination_redis_settings, @@ -658,10 +674,13 @@ from litellm.proxy.management_endpoints.management_v1.common import MANAGEMENT_V from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( router as model_access_group_management_router, ) -from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, - _add_team_model_to_db, - _deduplicate_litellm_router_models, +from litellm.proxy.management_endpoints.model_management_endpoints import ( # noqa: F401, RUF100 # legacy module exports + _add_model_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _add_team_model_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _deduplicate_litellm_router_models, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + add_model_to_db, + add_team_model_to_db, + deduplicate_litellm_router_models, live_model_ids_snapshot, ) from litellm.proxy.management_endpoints.model_management_endpoints import ( @@ -841,26 +860,33 @@ from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import ( from litellm.proxy.ui_crud_endpoints.user_banner_endpoints import ( router as user_banner_endpoints_router, ) -from litellm.proxy.utils import ( +from litellm.proxy.utils import ( # noqa: F401, RUF100 # legacy module exports PrismaClient, ProxyLogging, ProxyUpdateSpend, - _cache_user_row, - _get_docs_url, - _get_openapi_url, - _get_projected_spend_over_limit, - _get_redoc_url, - _is_projected_spend_over_limit, - _is_valid_team_configs, + _cache_user_row, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_docs_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_openapi_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_projected_spend_over_limit, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _get_redoc_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _is_projected_spend_over_limit, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _is_valid_team_configs, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + cache_user_row, evict_config_param, get_config_param, get_custom_url, + get_docs_url, get_error_message_str, + get_openapi_url, + get_projected_spend_over_limit, + get_redoc_url, get_server_root_path, handle_exception_on_proxy, hash_password, hash_token, invalidate_config_param, + is_projected_spend_over_limit, + is_valid_team_configs, litellm_config_cache, migrate_passwords_to_scrypt_async, model_dump_with_preserved_fields, @@ -1072,7 +1098,9 @@ custom_swagger_message: Final = ( ) ### CUSTOM BRANDING [ENTERPRISE FEATURE] ### -_title: Final = os.getenv("DOCS_TITLE", "LiteLLM API") if premium_user else "LiteLLM API" +title: Final = os.getenv("DOCS_TITLE", "LiteLLM API") if premium_user else "LiteLLM API" + +_title: Final = title _description: Final = ( os.getenv( "DOCS_DESCRIPTION", @@ -1238,7 +1266,7 @@ class _AiohttpConnectorKwargs(TypedDict, total=False): socket_factory: Callable[[_AiohttpAddrInfo], socket.socket] -async def _initialize_shared_aiohttp_session(): +async def initialize_shared_aiohttp_session() -> "ClientSession | None": """Initialize shared aiohttp session for connection reuse with connection limits.""" try: from aiohttp import ClientSession, DummyCookieJar, TCPConnector @@ -1278,6 +1306,9 @@ async def _initialize_shared_aiohttp_session(): return None +_initialize_shared_aiohttp_session: Final = initialize_shared_aiohttp_session + + async def _connect_to_count_stored_values() -> SupportsRawQueries: client: Final = prisma_client or PrismaClient( database_url=str(get_secret("DATABASE_URL")), proxy_logging_obj=proxy_logging_obj @@ -1428,10 +1459,12 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState # check if DATABASE_URL in environment - load from there if prisma_client is None: _db_url: Final[str | None] = get_secret("DATABASE_URL", None) - prisma_client = await ProxyStartupEvent._setup_prisma_client( - database_url=_db_url, - proxy_logging_obj=proxy_logging_obj, - user_api_key_cache=user_api_key_cache, + prisma_client = ( # rebind-ok: pre-existing rebinding on a rename-only line + await ProxyStartupEvent.setup_prisma_client( + database_url=_db_url, + proxy_logging_obj=proxy_logging_obj, + user_api_key_cache=user_api_key_cache, + ) ) await migrate_if_requested( @@ -1499,7 +1532,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState ## A coordination_redis block saved from the admin UI lives in the database, ## which is only reachable once the prisma client exists. Apply it here, before ## the coordination Redis is published to its consumers below. - db_coordination_redis_cache: Final = await ProxyStartupEvent._init_coordination_redis_from_db( + db_coordination_redis_cache: Final = await ProxyStartupEvent.init_coordination_redis_from_db( litellm_settings=proxy_config.get_config_state().get("litellm_settings") or {}, llm_router=llm_router, ) @@ -1510,11 +1543,11 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState ## when the proxy cache backend is not Redis ## transaction_buffer_redis_cache = redis_usage_cache if transaction_buffer_redis_cache is None: - transaction_buffer_redis_cache = ProxyStartupEvent._get_transaction_buffer_redis_cache( + transaction_buffer_redis_cache = ProxyStartupEvent.get_transaction_buffer_redis_cache( # rebind-ok: pre-existing rebinding on a rename-only line general_settings=general_settings ) - ProxyStartupEvent._initialize_startup_logging( + ProxyStartupEvent.initialize_startup_logging( llm_router=llm_router, proxy_logging_obj=proxy_logging_obj, redis_usage_cache=transaction_buffer_redis_cache, @@ -1553,12 +1586,12 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState verbose_proxy_logger.debug("Skipping OTel V2 provider setup: %s", e) ## Validate use_redis_transaction_buffer requires Redis cache ## - ProxyStartupEvent._validate_redis_transaction_buffer_config( + ProxyStartupEvent.validate_redis_transaction_buffer_config( general_settings=general_settings, redis_usage_cache=transaction_buffer_redis_cache, ) - ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings=general_settings) + ProxyStartupEvent.warn_if_mock_testing_params_enabled(general_settings=general_settings) ## SEMANTIC TOOL FILTER ## # Read litellm_settings from config for semantic filter initialization @@ -1567,7 +1600,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState _config: Final = proxy_config.get_config_state() _litellm_settings: Final = _config.get("litellm_settings", {}) verbose_proxy_logger.debug("litellm_settings keys = %s", list(_litellm_settings.keys())) - await ProxyStartupEvent._initialize_semantic_tool_filter( + await ProxyStartupEvent.initialize_semantic_tool_filter( llm_router=llm_router, litellm_settings=_litellm_settings, ) @@ -1576,28 +1609,28 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState verbose_proxy_logger.error("Semantic filter init failed: %s", e, exc_info=True) ## JWT AUTH ## - ProxyStartupEvent._initialize_jwt_auth( + ProxyStartupEvent.initialize_jwt_auth( general_settings=general_settings, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, ) - ProxyStartupEvent._attach_router_to_prompt_injection_detectors(llm_router=llm_router) + ProxyStartupEvent.attach_router_to_prompt_injection_detectors(llm_router=llm_router) verbose_proxy_logger.debug("prisma_client: %s", prisma_client) if prisma_client is not None and litellm.max_budget > 0: - ProxyStartupEvent._add_proxy_budget_to_db() + ProxyStartupEvent.add_proxy_budget_to_db() asyncio.create_task( - ProxyStartupEvent._warm_global_spend_cache( + ProxyStartupEvent.warm_global_spend_cache( user_api_key_cache=user_api_key_cache, prisma_client=prisma_client, ) ) - ProxyStartupEvent._warn_budget_without_db( + ProxyStartupEvent.warn_budget_without_db( max_budget=litellm.max_budget, prisma_client=prisma_client, ) - ProxyStartupEvent._warn_fail_closed_rate_limits_without_redis( + ProxyStartupEvent.warn_fail_closed_rate_limits_without_redis( fail_closed_rate_limit_enforcement=fail_closed_rate_limit_enforcement_enabled(general_settings), redis_usage_cache=redis_usage_cache, ) @@ -1616,10 +1649,10 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState else None ) if prisma_client is not None: - await ProxyStartupEvent._update_default_team_member_budget() + await ProxyStartupEvent.update_default_team_member_budget() ## SYNC UI SETTINGS ## - await ProxyStartupEvent._sync_ui_settings_to_general_settings() + await ProxyStartupEvent.sync_ui_settings_to_general_settings() # Start background health checks AFTER models are loaded and index is built if use_background_health_checks: @@ -1638,13 +1671,15 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState asyncio.create_task(_adaptive_router_flusher_loop()) ## [Optional] Initialize dd tracer - ProxyStartupEvent._init_dd_tracer() + ProxyStartupEvent.init_dd_tracer() ## [Optional] Initialize Pyroscope continuous profiling (env: LITELLM_ENABLE_PYROSCOPE=true) - ProxyStartupEvent._init_pyroscope() + ProxyStartupEvent.init_pyroscope() ## Initialize shared aiohttp session for connection reuse - shared_aiohttp_session = await _initialize_shared_aiohttp_session() + shared_aiohttp_session = ( # rebind-ok: pre-existing rebinding on a rename-only line + await initialize_shared_aiohttp_session() + ) model_info_refresh_disabled: Final = ( "disable_model_info_refresh" in general_settings and general_settings["disable_model_info_refresh"] is True @@ -1858,10 +1893,10 @@ def ensure_unique_openapi_operation_ids( app = FastAPI( - docs_url=_get_docs_url(), - redoc_url=_get_redoc_url(), - openapi_url=_get_openapi_url(), - title=_title, + docs_url=get_docs_url(), + redoc_url=get_redoc_url(), + openapi_url=get_openapi_url(), + title=title, description=_description, version=version, root_path=server_root_path, @@ -2689,7 +2724,7 @@ def mount_swagger_ui(): mount_swagger_ui() -docs_url: Final = _get_docs_url() +docs_url: Final = get_docs_url() root_redirect_url: Final[str | None] = os.getenv("ROOT_REDIRECT_URL") if docs_url != "/" and root_redirect_url is not None: @@ -2742,7 +2777,7 @@ user_api_key_cache: UserApiKeyCache = UserApiKeyCache( ) spend_counter_cache: Final = DualCache(default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value) cli_sso_session_cache: Final = DualCache(default_in_memory_ttl=CLI_SSO_SESSION_TTL_SECONDS) -model_max_budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=spend_counter_cache) +model_max_budget_limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=spend_counter_cache) litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter) redis_usage_cache: RedisCache | None = None # redis cache used for tracking spend, tpm/rpm limits polling_via_cache_enabled: Literal["all"] | list[str] | bool = False @@ -2779,7 +2814,7 @@ proxy_config_reload_interval_seconds = PROXY_CONFIG_RELOAD_INTERVAL_SECONDS litellm_master_key_hash = None disable_spend_logs = False jwt_handler: Final = JWTHandler() -prompt_injection_detection_obj: _OPTIONAL_PromptInjectionDetection | None = None +prompt_injection_detection_obj: Final[OPTIONAL_PromptInjectionDetection | None] = None store_model_in_db: bool = False open_telemetry_logger: OpenTelemetry | None = None ### GATEWAY REQUEST COUNTS (SGR) ### @@ -2794,7 +2829,7 @@ def _gateway_request_redis_buffer() -> GatewayRequestRedisBuffer | None: """Shares the spend writer's transaction-buffer Redis and pod lock when use_redis_transaction_buffer is on.""" writer: Final = proxy_logging_obj.db_spend_update_writer redis_cache: Final = writer.redis_update_buffer.redis_cache - if redis_cache is None or not writer.redis_update_buffer._should_commit_spend_updates_to_redis(): + if redis_cache is None or not writer.redis_update_buffer.should_commit_spend_updates_to_redis(): return None return GatewayRequestRedisBuffer(redis_cache=redis_cache, pod_lock_manager=writer.pod_lock_manager) @@ -2883,8 +2918,8 @@ def cost_tracking(): from litellm.integrations.shadow_eval_logger import ShadowEvalLogger spend_event_producer = build_spend_event_producer(CollectorSettings(), fallback=run_spend_event) - litellm.logging_callback_manager.add_litellm_callback(_ProxyDBLogger(spend_event_producer)) - litellm.logging_callback_manager.add_litellm_async_success_callback(_ProxyDBLogger(spend_event_producer)) + litellm.logging_callback_manager.add_litellm_callback(ProxyDBLogger(spend_event_producer)) + litellm.logging_callback_manager.add_litellm_async_success_callback(ProxyDBLogger(spend_event_producer)) litellm.logging_callback_manager.add_litellm_callback(ShadowEvalLogger()) @@ -3666,7 +3701,7 @@ async def _prepare_spend_counter_increment( 4. Increment is returned for the caller to apply via pipeline """ with service_target(SPEND_COUNTERS_TARGET): - await _ensure_spend_counter_initialized( + await ensure_spend_counter_initialized( counter_key=counter_key, source_cache_key=source_cache_key, ) @@ -3738,7 +3773,7 @@ async def _prepare_window_spend_counter_increment( return None with service_target(SPEND_COUNTERS_TARGET): - initialized: Final = await _ensure_window_spend_counter_initialized( + initialized: Final = await ensure_window_spend_counter_initialized( counter_key=counter_key, entity_type=entity_type, entity_id=entity_id, @@ -3750,10 +3785,10 @@ async def _prepare_window_spend_counter_increment( return PendingSpendIncrement(counter_key=counter_key, increment=increment) -async def _ensure_spend_counter_initialized( +async def ensure_spend_counter_initialized( counter_key: str, source_cache_key: str | list[str], -): +) -> None: is_warm: Final = await _is_spend_counter_cache_warm(counter_key=counter_key) if is_warm is False: # Shares the per-counter lock with get_current_spend. @@ -3767,7 +3802,10 @@ async def _ensure_spend_counter_initialized( # DB unavailable - fall back to in-process cache (may be stale). base_spend: Final = await _get_source_cache_base_spend(source_cache_key=source_cache_key) if base_spend > 0: - await _increment_spend_counter_cache(counter_key=counter_key, increment=base_spend) + await increment_spend_counter_cache(counter_key=counter_key, increment=base_spend) + + +_ensure_spend_counter_initialized: Final = ensure_spend_counter_initialized async def _get_source_cache_base_spend( @@ -3784,7 +3822,7 @@ async def _get_source_cache_base_spend( return 0.0 -async def _ensure_window_spend_counter_initialized( +async def ensure_window_spend_counter_initialized( counter_key: str, entity_type: str, entity_id: str, @@ -3813,6 +3851,9 @@ async def _ensure_window_spend_counter_initialized( return True +_ensure_window_spend_counter_initialized: Final = ensure_window_spend_counter_initialized + + @with_service_target(SPEND_COUNTERS_TARGET) async def _is_spend_counter_cache_warm(counter_key: str) -> bool: batched: Final = await read_batched_spend_counter(counter_key) @@ -3849,7 +3890,7 @@ async def increment_spend_counter(counter_key: str, increment: float): """Public raw-counter increment for budget domains outside the entity scopes (e.g. shadow eval's per-leg spend), sharing the primitive the entity counters use so invalidation and read semantics can never drift.""" - return await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) + return await increment_spend_counter_cache(counter_key=counter_key, increment=increment) @with_service_target(SPEND_COUNTERS_TARGET) @@ -3864,7 +3905,7 @@ async def refresh_spend_counter_ttl(counter_key: str) -> bool: @with_service_target(SPEND_COUNTERS_TARGET) -async def _increment_spend_counter_cache(counter_key: str, increment: float): +async def increment_spend_counter_cache(counter_key: str, increment: float) -> float | None: if spend_counter_cache.redis_cache is not None: try: current_value: Final = await spend_counter_cache.redis_cache.async_increment( @@ -3873,7 +3914,7 @@ async def _increment_spend_counter_cache(counter_key: str, increment: float): refresh_ttl=True, ) except Exception: - await _invalidate_spend_counter(counter_key=counter_key) + await invalidate_spend_counter(counter_key=counter_key) raise spend_counter_cache.in_memory_cache.set_cache( key=counter_key, @@ -3887,8 +3928,11 @@ async def _increment_spend_counter_cache(counter_key: str, increment: float): ) +_increment_spend_counter_cache: Final = increment_spend_counter_cache + + @with_service_target(SPEND_COUNTERS_TARGET) -async def _invalidate_spend_counter(counter_key: str): +async def invalidate_spend_counter(counter_key: str) -> None: forget_spend_counter(counter_key) spend_counter_cache.in_memory_cache.delete_cache(key=counter_key) if spend_counter_cache.redis_cache is not None: @@ -3902,6 +3946,9 @@ async def _invalidate_spend_counter(counter_key: str): ) +_invalidate_spend_counter: Final = invalidate_spend_counter + + async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) -> None: if _defer_spend_counter_increments(pending): return @@ -3944,7 +3991,7 @@ def _settle_spend_counter_increment(item: PendingSpendIncrement) -> Callable[[as verbose_proxy_logger.warning( "Spend counter %s increment did not land in the post-call pipeline; invalidating it", item.counter_key ) - await _invalidate_spend_counter(counter_key=item.counter_key) + await invalidate_spend_counter(counter_key=item.counter_key) return settle @@ -3957,7 +4004,7 @@ async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrem try: return await run_spend_counter_pipeline(pending=pending) except Exception: - await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending)) + await asyncio.gather(*(invalidate_spend_counter(counter_key=item.counter_key) for item in pending)) raise @@ -4106,14 +4153,14 @@ async def update_cache( existing_spend_obj.soft_budget_cooldown is False and existing_spend_obj.soft_budget is not None and ( - _is_projected_spend_over_limit( + is_projected_spend_over_limit( current_spend=new_spend, soft_budget_limit=existing_spend_obj.soft_budget, ) is True ) ): - projected_spend, projected_exceeded_date = _get_projected_spend_over_limit( + projected_spend, projected_exceeded_date = get_projected_spend_over_limit( current_spend=new_spend, soft_budget_limit=existing_spend_obj.soft_budget, ) @@ -4491,13 +4538,13 @@ def _schedule_background_health_check_db_save( import time as time_module from litellm.proxy.health_endpoints._health_endpoints import ( - _save_background_health_checks_to_db, + save_background_health_checks_to_db, ) checked_by: Final = shared_health_manager.pod_id if shared_health_manager is not None else "background_health_check" start_time: Final = time_module.time() save: Final = partial( - _save_background_health_checks_to_db, + save_background_health_checks_to_db, prisma_client, model_list, healthy_endpoints, @@ -5016,7 +5063,7 @@ def _resolve_coordination_redis_env_refs(raw_params: Mapping[str, object]) -> di } -def _build_redis_usage_cache(redis_params: Mapping[str, object]) -> RedisCache: +def build_redis_usage_cache(redis_params: Mapping[str, object]) -> RedisCache: """ Builds the proxy's coordination Redis client from resolved connection params. Cluster-mode targets (explicit `startup_nodes` or the @@ -5035,7 +5082,10 @@ def _build_redis_usage_cache(redis_params: Mapping[str, object]) -> RedisCache: return RedisCache(**non_node_params) -def _environment_has_redis_connection_target() -> bool: +_build_redis_usage_cache: Final = build_redis_usage_cache + + +def environment_has_redis_connection_target() -> bool: """ Whether the REDIS_* environment variables name a Redis to connect to (host, url, cluster nodes, or sentinel nodes). Read-only: callers that only need to @@ -5051,6 +5101,9 @@ def _environment_has_redis_connection_target() -> bool: ) +_environment_has_redis_connection_target: Final = environment_has_redis_connection_target + + def _build_redis_usage_cache_from_environment() -> RedisCache | None: """ Builds a standalone coordination Redis from REDIS_* environment variables. @@ -5062,9 +5115,9 @@ def _build_redis_usage_cache_from_environment() -> RedisCache | None: Returns None when the environment carries no connection target (host, url, cluster nodes, or sentinel nodes). """ - if not _environment_has_redis_connection_target(): + if not environment_has_redis_connection_target(): return None - return _build_redis_usage_cache(litellm._redis._redis_kwargs_from_environment()) + return build_redis_usage_cache(litellm._redis._redis_kwargs_from_environment()) def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache: bool) -> None: @@ -5714,7 +5767,7 @@ class ProxyConfig: environment_variables: Final = new_config.get("environment_variables") if include_env_vars and environment_variables is not None: encrypted_environment_variables: Final = ( - self._encrypt_env_variables_for_db(environment_variables=environment_variables) + self.encrypt_env_variables_for_db(environment_variables=environment_variables) if isinstance(environment_variables, dict) and environment_variables else environment_variables ) @@ -5875,7 +5928,7 @@ class ProxyConfig: existing: Final[dict] = dict(row.param_value) if row is not None and row.param_value is not None else {} to_set: Final = {k: v for k, v in updates.items() if v is not None} - encrypted: Final = self._encrypt_env_variables_for_db(environment_variables=to_set) if to_set else {} + encrypted: Final = self.encrypt_env_variables_for_db(environment_variables=to_set) if to_set else {} deleted_keys: Final = {k for k, v in updates.items() if v is None} merged: Final = {**{k: v for k, v in existing.items() if k not in deleted_keys}, **encrypted} @@ -6028,7 +6081,7 @@ class ProxyConfig: "set one of host, url, startup_nodes, or sentinel_nodes" ) - coordination_redis_cache: Final = _build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) + coordination_redis_cache: Final = build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) _attach_redis_usage_cache( coordination_redis_cache, enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, @@ -6096,7 +6149,7 @@ class ProxyConfig: ) return env_coordination_redis_cache - def _init_cache( + def init_cache( self, cache_params: dict, enable_redis_auth_cache: bool = False, @@ -6142,6 +6195,8 @@ class ProxyConfig: verbose_proxy_logger.info("litellm_config_cache: no Redis configured; cluster-wide cache sharing disabled.") return resolved_usage_cache + _init_cache = init_cache + def switch_on_llm_response_caching(self): """ Enable caching on the router by setting cache_responses=True. @@ -6512,7 +6567,7 @@ class ProxyConfig: ## to pass a complete url, or set ssl=True, etc. just set it as `os.environ[REDIS_URL] = `, _redis.py checks for REDIS specific environment variables _set_redis_usage_cache( - self._init_cache( + self.init_cache( cache_params=cache_params, enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, ) @@ -7731,7 +7786,7 @@ class ProxyConfig: displaced=previous.displaced + _entries_missing_from(before, after), ) - def _encrypt_env_variables(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: + def encrypt_env_variables(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: """ Encrypts a dictionary of environment variables and returns them. """ @@ -7741,7 +7796,9 @@ class ProxyConfig: encrypted_env_vars[k] = encrypted_value return encrypted_env_vars - def _decrypt_and_set_db_env_variables( + _encrypt_env_variables = encrypt_env_variables + + def decrypt_and_set_db_env_variables( self, environment_variables: dict, return_original_value: bool = False ) -> dict: """ @@ -7770,7 +7827,9 @@ class ProxyConfig: verbose_proxy_logger.error("Error setting env variable: %s - %s", k, str(e)) return decrypted_env_vars - def _decrypt_db_variables(self, variables_dict: dict) -> dict: + _decrypt_and_set_db_env_variables = decrypt_and_set_db_env_variables + + def decrypt_db_variables(self, variables_dict: dict) -> dict: """ Decrypts a dictionary of variables and returns them. """ @@ -7780,7 +7839,9 @@ class ProxyConfig: decrypted_variables[k] = decrypted_value return decrypted_variables - def _encrypt_env_variables_for_db(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: + _decrypt_db_variables = decrypt_db_variables + + def encrypt_env_variables_for_db(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: """ Idempotently encrypt environment variables for a DB write. @@ -7795,12 +7856,14 @@ class ProxyConfig: _decrypt_and_set_db_env_variables): this is a write path, and loading values into os.environ is the read path's responsibility. """ - decrypted_env_vars: Final = self._decrypt_db_variables(environment_variables) - return self._encrypt_env_variables( + decrypted_env_vars: Final = self.decrypt_db_variables(environment_variables) + return self.encrypt_env_variables( environment_variables=decrypted_env_vars, new_encryption_key=new_encryption_key, ) + _encrypt_env_variables_for_db = encrypt_env_variables_for_db + @staticmethod def _parse_router_settings_value(value: object) -> dict | None: """ @@ -8055,7 +8118,7 @@ class ProxyConfig: async def _apply_pass_through_settings(self, db_values: Mapping[str, SettingsJsonValue]) -> None: db_endpoints: Final = db_values.get("pass_through_endpoints") if isinstance(db_endpoints, list): - await self._serve_pass_through_endpoints(db_endpoints) + await self.serve_pass_through_endpoints(db_endpoints) return if "pass_through_endpoints" not in self.settings: self._publish_pass_through_endpoints(()) @@ -8065,10 +8128,12 @@ class ProxyConfig: list(db_endpoints), config_passthrough_endpoints ) - async def _serve_pass_through_endpoints(self, db_endpoints: Sequence[SettingsJsonValue]) -> None: + async def serve_pass_through_endpoints(self, db_endpoints: Sequence[SettingsJsonValue]) -> None: self._publish_pass_through_endpoints(db_endpoints) await initialize_pass_through_endpoints(pass_through_endpoints=_ENDPOINT_DICTS.validate_python(db_endpoints)) + _serve_pass_through_endpoints = serve_pass_through_endpoints + async def _apply_boolean_settings(self, db_values: Mapping[str, SettingsJsonValue]) -> None: for key in ( "store_prompts_in_spend_logs", @@ -8200,7 +8265,7 @@ class ProxyConfig: def _prepared_db_settings_values(self, section: Section, value: object) -> Mapping[str, SettingsJsonValue]: if section == "environment_variables": - decrypted: Final = self._decrypt_and_set_db_env_variables( + decrypted: Final = self.decrypt_and_set_db_env_variables( dict(_as_settings_mapping(value)), return_original_value=True ) normalized: Final = { @@ -8229,7 +8294,7 @@ class ProxyConfig: row for row in self.auto_router_db_catalog if row.model_id not in model_ids ) - async def _get_models_from_db(self, prisma_client: PrismaClient) -> Sequence[_ProxyModelRow] | None: + async def get_models_from_db(self, prisma_client: PrismaClient) -> Sequence[_ProxyModelRow] | None: """ Fetch all model deployments from the DB. @@ -8257,6 +8322,8 @@ class ProxyConfig: ) return None + _get_models_from_db = get_models_from_db + async def add_deployment( self, prisma_client: PrismaClient, @@ -8289,9 +8356,9 @@ class ProxyConfig: await sync_ui_settings_to_general_settings(prisma_client) async with MODEL_RECONCILE_LOCK: - return await self._add_deployment_locked(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj) + return await self.add_deployment_locked(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj) - async def _add_deployment_locked( + async def add_deployment_locked( self, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging, @@ -8318,7 +8385,7 @@ class ProxyConfig: ) load_models: Final = self._should_load_db_object(object_type="models") - new_models: Final = await self._get_models_from_db(prisma_client=prisma_client) if load_models else None + new_models: Final = await self.get_models_from_db(prisma_client=prisma_client) if load_models else None await self.get_credentials(prisma_client=prisma_client) if load_models: still_desired_ids = await self._update_llm_router( @@ -8348,6 +8415,8 @@ class ProxyConfig: live_after=None if still_desired_ids is None else live_model_ids_snapshot(), ) + _add_deployment_locked = add_deployment_locked + def start_config_sync_subscriber( self, prisma_client: PrismaClient, @@ -8451,7 +8520,7 @@ class ProxyConfig: await CacheSettingsManager.init_cache_settings_in_db(prisma_client=prisma_client, proxy_config=self) if self._should_load_db_object(object_type="semantic_filter_settings"): - await self._init_semantic_filter_settings_in_db(prisma_client=prisma_client) + await self.init_semantic_filter_settings_in_db(prisma_client=prisma_client) if self._should_load_db_object(object_type=SupportedDBObjectType.WEBSEARCH_INTERCEPTION_SETTINGS): await self.init_websearch_interception_settings_in_db(prisma_client=prisma_client) @@ -8470,7 +8539,7 @@ class ProxyConfig: db_values: Final = self._prepared_db_settings_values("litellm_settings", raw_settings) self._apply_litellm_settings_db_values(db_values) - async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient): + async def init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient) -> None: """ Initialize MCP semantic filter settings from database. Called periodically (approximately every 10 seconds) by background task to hot-reload settings across all pods. @@ -8534,6 +8603,8 @@ class ProxyConfig: except Exception as e: verbose_proxy_logger.exception("Error initializing semantic filter settings from DB: %s", e) + _init_semantic_filter_settings_in_db = init_semantic_filter_settings_in_db + async def init_websearch_interception_settings_in_db(self, prisma_client: PrismaClient): """ Initialize web search interception settings from database. @@ -8614,7 +8685,7 @@ class ProxyConfig: sso_settings.sso_settings.pop("team_mappings", None) sso_settings.sso_settings.pop("ui_access_mode", None) uppercase_sso_settings: Final = {key.upper(): value for key, value in sso_settings.sso_settings.items()} - self._decrypt_and_set_db_env_variables(environment_variables=uppercase_sso_settings) + self.decrypt_and_set_db_env_variables(environment_variables=uppercase_sso_settings) except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.py::ProxyConfig:_init_sso_settings_in_db - %s", e @@ -8628,10 +8699,10 @@ class ProxyConfig: """ from litellm.proxy.management_endpoints.config_override_endpoints import ( HASHICORP_ENV_VAR_MAPPING, - _clear_hashicorp_vault_state, - _get_current_env_values, - _parse_config_value, - _set_env_vars, + clear_hashicorp_vault_state, + get_current_env_values, + parse_config_value, + set_env_vars, ) try: @@ -8648,28 +8719,28 @@ class ProxyConfig: if db_record is None or db_record.config_value is None: if self._last_hashicorp_vault_config is not None: - _clear_hashicorp_vault_state(self) + clear_hashicorp_vault_state(self) return - config_data: Final = _parse_config_value(db_record.config_value) + config_data: Final = parse_config_value(db_record.config_value) # Skip reinit if config hasn't changed since last poll if self._last_hashicorp_vault_config == config_data: return # Decrypt all fields and set env vars - decrypted_data: Final = self._decrypt_db_variables(config_data) + decrypted_data: Final = self.decrypt_db_variables(config_data) # Snapshot current env vars so we can restore on failure - previous_env: Final = _get_current_env_values(HASHICORP_ENV_VAR_MAPPING) - _set_env_vars(decrypted_data) + previous_env: Final = get_current_env_values(HASHICORP_ENV_VAR_MAPPING) + set_env_vars(decrypted_data) # Reinitialize the secret manager try: self.initialize_secret_manager(key_management_system="hashicorp_vault") except Exception: # Restore previous working env vars instead of wiping all - _set_env_vars(previous_env) + set_env_vars(previous_env) raise self._last_hashicorp_vault_config = config_data.copy() @@ -8689,10 +8760,10 @@ class ProxyConfig: from litellm.proxy.management_endpoints.config_override_endpoints import ( CYBERARK_ENV_VAR_MAPPING, _clear_cyberark_state, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module - _get_current_env_values, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module - _parse_config_value, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module - _set_env_vars, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module _snapshot_cyberark_boot_env, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module + get_current_env_values, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module + parse_config_value, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module + set_env_vars, # pyright: ignore[reportPrivateUsage] # module-internal helper shared with the endpoint module ) try: @@ -8712,22 +8783,22 @@ class ProxyConfig: _clear_cyberark_state(self) return - config_data: Final = _parse_config_value(db_record.config_value) + config_data: Final = parse_config_value(db_record.config_value) # Skip reinit if config hasn't changed since last poll if self._last_cyberark_config == config_data: return - decrypted_data: Final = self._decrypt_db_variables(config_data) + decrypted_data: Final = self.decrypt_db_variables(config_data) _snapshot_cyberark_boot_env(self) - previous_env: Final = _get_current_env_values(CYBERARK_ENV_VAR_MAPPING) - _set_env_vars(decrypted_data, CYBERARK_ENV_VAR_MAPPING) + previous_env: Final = get_current_env_values(CYBERARK_ENV_VAR_MAPPING) + set_env_vars(decrypted_data, CYBERARK_ENV_VAR_MAPPING) try: self.initialize_secret_manager(key_management_system="cyberark") except Exception: - _set_env_vars(previous_env, CYBERARK_ENV_VAR_MAPPING) + set_env_vars(previous_env, CYBERARK_ENV_VAR_MAPPING) raise self._last_cyberark_config = config_data.copy() @@ -9583,7 +9654,7 @@ def _restamp_streaming_chunk_model( fallback_was_attempted: bool = False, fallback_model_from_metadata: str | None = None, ) -> tuple[Any, bool]: - if _should_return_raw_model_name(request_data): + if should_return_raw_model_name(request_data): return chunk, model_mismatch_logged target_model: Final = fallback_model_from_metadata if fallback_was_attempted else requested_model_from_client @@ -9599,7 +9670,7 @@ def _restamp_streaming_chunk_model( return chunk, model_mismatch_logged # For Azure Model Router, preserve the actual model used in each chunk - if not fallback_was_attempted and _is_azure_model_router_request(requested_model_from_client): + if not fallback_was_attempted and is_azure_model_router_request(requested_model_from_client): return chunk, model_mismatch_logged # For fastest_response batch completions, preserve the winning model's name @@ -10179,7 +10250,7 @@ async def async_data_generator( # The iterator-wrap path fires deferred logging itself; fire it # here for the no-wrap fast path so non-callback deployments # still flush their post-stream logging. - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) if raw_sse_buffer: yield (raw_sse_buffer if raw_sse_buffer.endswith(_SSE_FRAME_DELIMITERS) else raw_sse_buffer + "\n\n") @@ -10244,7 +10315,7 @@ async def async_data_generator( stream_completed = True yield f"data: {error_returned}\n\n" finally: - await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( + await ProxyBaseLLMRequestProcessing.finalize_streaming_generator_cleanup( request=request, request_data=request_data, response=response, @@ -10318,20 +10389,20 @@ def giveup(e): class ProxyStartupEvent: @staticmethod - def _attach_router_to_prompt_injection_detectors(llm_router: Router | None) -> None: - for callback in litellm.logging_callback_manager.get_custom_loggers_for_type( - _OPTIONAL_PromptInjectionDetection - ): - if isinstance(callback, _OPTIONAL_PromptInjectionDetection): + def attach_router_to_prompt_injection_detectors(llm_router: Router | None) -> None: + for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(OPTIONAL_PromptInjectionDetection): + if isinstance(callback, OPTIONAL_PromptInjectionDetection): callback.update_environment(router=llm_router) + _attach_router_to_prompt_injection_detectors = attach_router_to_prompt_injection_detectors + @staticmethod async def refresh_model_info() -> None: if llm_router is not None: await llm_router.arefresh_model_info() @staticmethod - def _warn_budget_without_db(max_budget: float | None, prisma_client: PrismaClient | None) -> None: + def warn_budget_without_db(max_budget: float | None, prisma_client: PrismaClient | None) -> None: if prisma_client is not None or not max_budget or max_budget <= 0: return @@ -10343,8 +10414,10 @@ class ProxyStartupEvent: max_budget, ) + _warn_budget_without_db = warn_budget_without_db + @staticmethod - def _warn_fail_closed_rate_limits_without_redis( + def warn_fail_closed_rate_limits_without_redis( fail_closed_rate_limit_enforcement: bool, redis_usage_cache: RedisCache | None ) -> None: if redis_usage_cache is not None or not fail_closed_rate_limit_enforcement: @@ -10357,8 +10430,10 @@ class ProxyStartupEvent: "across pods and make the setting effective." ) + _warn_fail_closed_rate_limits_without_redis = warn_fail_closed_rate_limits_without_redis + @classmethod - def _initialize_startup_logging( + def initialize_startup_logging( cls, llm_router: Router | None, proxy_logging_obj: ProxyLogging, @@ -10370,8 +10445,10 @@ class ProxyStartupEvent: proxy_logging_obj.startup_event(llm_router=llm_router, redis_usage_cache=redis_usage_cache) + _initialize_startup_logging = initialize_startup_logging + @staticmethod - def _warn_if_mock_testing_params_enabled(general_settings: dict) -> None: + def warn_if_mock_testing_params_enabled(general_settings: dict) -> None: """Announce, loudly, that any caller may inject synthetic failures.""" from litellm.proxy.route_llm_request import ( GATED_MOCK_PARAM_NAMES, @@ -10401,8 +10478,10 @@ class ProxyStartupEvent: "=" * 72, ) + _warn_if_mock_testing_params_enabled = warn_if_mock_testing_params_enabled + @staticmethod - def _validate_redis_transaction_buffer_config( + def validate_redis_transaction_buffer_config( general_settings: dict, redis_usage_cache: RedisCache | None, ): @@ -10431,8 +10510,10 @@ class ProxyStartupEvent: "Redis for the transaction buffer." ) + _validate_redis_transaction_buffer_config = validate_redis_transaction_buffer_config + @staticmethod - async def _init_coordination_redis_from_db( + async def init_coordination_redis_from_db( litellm_settings: Mapping[str, object], llm_router: Router | None, ) -> RedisCache | None: @@ -10460,7 +10541,7 @@ class ProxyStartupEvent: ) return None - coordination_redis_cache: Final = _build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) + coordination_redis_cache: Final = build_redis_usage_cache(coordination_params.model_dump(exclude_none=True)) _attach_redis_usage_cache( coordination_redis_cache, enable_redis_auth_cache=litellm_settings.get("enable_redis_auth_cache", False) is True, @@ -10473,8 +10554,10 @@ class ProxyStartupEvent: ) return coordination_redis_cache + _init_coordination_redis_from_db = init_coordination_redis_from_db + @staticmethod - def _get_transaction_buffer_redis_cache( + def get_transaction_buffer_redis_cache( general_settings: dict, ) -> RedisCache | None: """ @@ -10496,8 +10579,10 @@ class ProxyStartupEvent: return _build_redis_usage_cache_from_environment() + _get_transaction_buffer_redis_cache = get_transaction_buffer_redis_cache + @classmethod - async def _initialize_semantic_tool_filter( + async def initialize_semantic_tool_filter( cls, llm_router: Router | None, litellm_settings: dict[str, Any], @@ -10529,8 +10614,10 @@ class ProxyStartupEvent: # Only warn if the feature was configured but failed to initialize verbose_proxy_logger.warning("Semantic tool filter hook was configured but failed to initialize") + _initialize_semantic_tool_filter = initialize_semantic_tool_filter + @classmethod - def _initialize_jwt_auth( + def initialize_jwt_auth( cls, general_settings: dict, prisma_client: PrismaClient | None, @@ -10563,14 +10650,18 @@ class ProxyStartupEvent: jwt_handler.bind_agent_lookup(global_agent_registry) + _initialize_jwt_auth = initialize_jwt_auth + @classmethod - def _add_proxy_budget_to_db(cls): + def add_proxy_budget_to_db(cls): """Adds a global proxy budget to db""" if litellm.budget_duration is None: raise Exception("budget_duration not set on Proxy. budget_duration is required to use max_budget.") asyncio.create_task(cls._upsert_proxy_budget_with_reset_at_backfill()) + _add_proxy_budget_to_db = add_proxy_budget_to_db + @classmethod async def _upsert_proxy_budget_with_reset_at_backfill(cls) -> None: """ @@ -10625,7 +10716,7 @@ class ProxyStartupEvent: verbose_proxy_logger.warning("Failed to backfill budget_reset_at on proxy admin row: %s", e) @classmethod - async def _warm_global_spend_cache( + async def warm_global_spend_cache( cls, user_api_key_cache: UserApiKeyCache, prisma_client: PrismaClient, @@ -10633,7 +10724,7 @@ class ProxyStartupEvent: """Warm global spend cache once at startup to reduce impact of first wave of requests.""" try: cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY - await _fetch_global_spend_with_event_coordination( + await fetch_global_spend_with_event_coordination( cache_key=cache_key, user_api_key_cache=user_api_key_cache, prisma_client=prisma_client, @@ -10641,8 +10732,10 @@ class ProxyStartupEvent: except Exception as e: verbose_proxy_logger.debug("Global spend cache warm-up at startup skipped or failed: %s", e) + _warm_global_spend_cache = warm_global_spend_cache + @classmethod - async def _update_default_team_member_budget(cls): + async def update_default_team_member_budget(cls): """Update the default team member budget""" if litellm.default_internal_user_params is None: return @@ -10659,8 +10752,10 @@ class ProxyStartupEvent: user_api_key_dict=UserAPIKeyAuth(token=hash_token(master_key)), ) + _update_default_team_member_budget = update_default_team_member_budget + @classmethod - async def _sync_ui_settings_to_general_settings(cls): + async def sync_ui_settings_to_general_settings(cls): """Apply the persisted UI settings to general_settings before this pod serves traffic.""" if prisma_client is None: return @@ -10668,6 +10763,8 @@ class ProxyStartupEvent: if applied: verbose_proxy_logger.info("Synced UI settings to general_settings on startup: %s", list(applied)) + _sync_ui_settings_to_general_settings = sync_ui_settings_to_general_settings + @classmethod async def _load_heuristic_v1_tuning_baselines( cls, prisma_client: PrismaClient, deployments: Sequence[Mapping[str, object]] @@ -10720,7 +10817,7 @@ class ProxyStartupEvent: cls, prisma_client: PrismaClient, llm_router: Router | None, limit: int | None ) -> Mapping[str, str] | None: """Load a complete baseline and reject a startup that exceeds the tuning quota.""" - db_models: Final = await proxy_config._get_models_from_db(prisma_client) + db_models: Final = await proxy_config.get_models_from_db(prisma_client) if db_models is None: verbose_proxy_logger.warning("Heuristic-v1 tuning baseline unavailable, gate not enforced this boot") return None @@ -10864,10 +10961,10 @@ class ProxyStartupEvent: ### MONITOR SPEND LOGS QUEUE (queue-size-based job) ### if general_settings.get("disable_spend_logs", False) is False: - from litellm.proxy.utils import _monitor_spend_logs_queue + from litellm.proxy.utils import monitor_spend_logs_queue monitor_task: Final = asyncio.create_task( - _monitor_spend_logs_queue( + monitor_spend_logs_queue( prisma_client=prisma_client, db_writer_client=db_writer_client, proxy_logging_obj=proxy_logging_obj, @@ -11520,7 +11617,7 @@ class ProxyStartupEvent: await _scheduled_fallback_stats() @classmethod - async def _setup_prisma_client( + async def setup_prisma_client( cls, database_url: str | None, proxy_logging_obj: ProxyLogging, @@ -11574,8 +11671,10 @@ class ProxyStartupEvent: ) return connected_client + _setup_prisma_client = setup_prisma_client + @classmethod - def _init_dd_tracer(cls): + def init_dd_tracer(cls): """ Initialize dd tracer - if `USE_DDTRACE=true` in .env @@ -11599,8 +11698,10 @@ class ProxyStartupEvent: prof.start() verbose_proxy_logger.debug("Datadog Profiler started......") + _init_dd_tracer = init_dd_tracer + @classmethod - def _init_pyroscope(cls): + def init_pyroscope(cls): """ Optional continuous profiling via Grafana Pyroscope. @@ -11682,6 +11783,8 @@ class ProxyStartupEvent: "Pyroscope profiling will not run. Install with: pip install pyroscope-io" ) + _init_pyroscope = init_pyroscope + #### API ENDPOINTS #### async def _names_hidden_by_listing_callbacks( @@ -11774,7 +11877,7 @@ async def model_list( create_anthropic_model_list_response, ) from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_privileges, + user_has_admin_privileges, ) from litellm.proxy.utils import ( create_model_info_response, @@ -11814,7 +11917,9 @@ async def model_list( # Check if scope=expand is requested and user has admin privileges should_expand_scope = False if scope == "expand": - should_expand_scope = _user_has_admin_view(user_api_key_dict) or await _user_has_admin_privileges( + should_expand_scope = user_api_key_has_admin_view( # rebind-ok: pre-existing rebinding on a rename-only line + user_api_key_dict + ) or await user_has_admin_privileges( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -12176,7 +12281,7 @@ async def chat_completion( """ global general_settings, user_debug, proxy_logging_obj, llm_model_list global user_temperature, user_request_timeout, user_max_tokens, user_api_base - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) if user_api_key_dict is not None: if not isinstance(data.get("metadata"), dict): # Covers both missing and JSON-string metadata (multipart / @@ -12291,7 +12396,7 @@ async def chat_completion( _chat_response.usage = _usage return _chat_response except Exception as e: - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -12337,7 +12442,7 @@ async def completion( global user_temperature, user_request_timeout, user_max_tokens, user_api_base data = {} try: - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line if user_api_key_dict is not None: if data.get("metadata") is None: data["metadata"] = {} @@ -12518,7 +12623,7 @@ async def embeddings( """ global proxy_logging_obj - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: ### HANDLE TOKEN ARRAY INPUT DECODING ### @@ -12585,7 +12690,7 @@ async def embeddings( return response except Exception as e: - raise await base_llm_response_processor._handle_llm_api_exception( + raise await base_llm_response_processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -13167,10 +13272,10 @@ async def realtime_websocket_endpoint( request: Final = Request(scope=scope) - request._url = websocket.url + request._url = websocket.url # pyright: ignore[reportPrivateUsage] # Starlette WebSocket URL storage async def return_body(): - return _realtime_request_body(route_model) + return realtime_request_body(route_model) request.body = return_body @@ -14445,7 +14550,7 @@ async def non_admin_all_models( ) # de-duplicate models. Only return unique model ids - unique_models: Final = _deduplicate_litellm_router_models(models=all_models) + unique_models: Final = deduplicate_litellm_router_models(models=all_models) return unique_models @@ -14673,7 +14778,7 @@ async def _populate_team_access_on_models( """ user_teams: list[str] | Literal["*"] | None = None direct_access_models: Sequence[str] = () - if _user_has_admin_view(user_api_key_dict): + if user_api_key_has_admin_view(user_api_key_dict): user_teams = "*" direct_access_models = tuple(llm_router.get_model_ids(exclude_team_models=True)) # access to all models elif user_api_key_dict.user_id is not None: @@ -16515,7 +16620,7 @@ async def model_deprecations( return collect_model_deprecations(llm_router=llm_router, warn_within_days=warn_within_days) -def _get_model_group_info( +def get_model_group_info( llm_router: Router, all_models_str: Sequence[str], model_group: str | None ) -> list[ModelGroupInfoProxy]: model_groups: Final[list[ModelGroupInfoProxy]] = [] @@ -16549,6 +16654,9 @@ def _get_model_group_info( return model_groups +_get_model_group_info: Final = get_model_group_info + + @router.get( "/model_group/info", tags=["model management"], @@ -16757,7 +16865,7 @@ async def model_group_info( ) model_groups: Final = await append_agents_to_model_group( - model_groups=_get_model_group_info( + model_groups=get_model_group_info( llm_router=llm_router, all_models_str=listed_group_names, model_group=model_group ), user_api_key_dict=user_api_key_dict, @@ -16854,7 +16962,7 @@ async def alerting_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, @@ -16982,7 +17090,7 @@ async def async_queue_request( data["proxy_server_request"] = { "url": str(request.url), "method": request.method, - "headers": _safe_get_request_headers(request).copy(), + "headers": safe_get_request_headers(request).copy(), "body": copy.copy(data), # use copy instead of deepcopy } @@ -17007,7 +17115,7 @@ async def async_queue_request( data["metadata"]["user_api_key"] = logged_api_key data["metadata"]["user_api_key_hash"] = logged_api_key data["metadata"]["user_api_key_metadata"] = strip_callback_config(user_api_key_dict.metadata) - _headers: Final = _safe_get_request_headers(request).copy() + _headers: Final = safe_get_request_headers(request).copy() _headers.pop("authorization", None) # do not store the original `sk-..` api key in the db data["metadata"]["headers"] = _headers data["metadata"]["user_api_key_alias"] = getattr(user_api_key_dict, "key_alias", None) @@ -17161,8 +17269,8 @@ async def login(request: Request): # _is_same_origin_return_path (strictly relative path) so it can never be an open redirect, and the # one-shot cookie is cleared after use. from litellm.proxy.management_endpoints.ui_sso import ( - _sso_return_to_redirect, set_session_token_cookie, + sso_return_to_redirect, ) # Resume through the SAME resumer the SSO callback uses, rather than a second, narrower arm. @@ -17175,7 +17283,7 @@ async def login(request: Request): cp_return_to: Final = request.cookies.get("litellm_cp_return_to") if cp_return_to: try: - resumed = await _sso_return_to_redirect( + resumed = await sso_return_to_redirect( # rebind-ok: pre-existing rebinding on a rename-only line return_to=cp_return_to, jwt_token=jwt_token, redis_usage_cache=redis_usage_cache, @@ -17935,11 +18043,14 @@ async def new_invitation(data: InvitationNew, user_api_key_dict: UserAPIKeyAuth ) # Allow proxy admins and org/team admins (admin status from DB via get_user_object) - has_access = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN or await _user_has_admin_privileges( - user_api_key_dict=user_api_key_dict, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, + has_access: Final = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or await user_has_admin_privileges( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) ) if not has_access: raise HTTPException( @@ -18000,7 +18111,7 @@ async def invitation_info(invitation_id: str, user_api_key_dict: UserAPIKeyAuth detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, @@ -18107,7 +18218,7 @@ async def invitation_delete( # Proxy admins can delete any invitation; org admins only their own is_proxy_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN - is_other_admin: Final = await _user_has_admin_privileges( + is_other_admin: Final = await user_has_admin_privileges( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, @@ -18293,7 +18404,7 @@ async def update_config( existing = await _read_section("environment_variables") before_environment_variables: Final = copy.deepcopy(existing) existing.update( - proxy_config._encrypt_env_variables_for_db(environment_variables=config_info.environment_variables) + proxy_config.encrypt_env_variables_for_db(environment_variables=config_info.environment_variables) ) await _upsert_section("environment_variables", existing) asyncio.create_task( @@ -18557,7 +18668,7 @@ async def update_config_general_settings( proxy_config.settings.apply_db_row("general_settings", general_settings) if is_resource_list("general_settings", data.field_name): stored_endpoints: Final = general_settings.get("pass_through_endpoints") - await proxy_config._serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) + await proxy_config.serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) asyncio.create_task( create_config_audit_log( "general_settings", "updated", before_general_settings, general_settings, user_api_key_dict @@ -18748,7 +18859,7 @@ async def get_config_general_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": CommonProxyErrors.not_allowed_access.value}, @@ -18953,7 +19064,7 @@ async def get_config_list( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": f"{CommonProxyErrors.not_allowed_access.value}, your role={user_api_key_dict.user_role}"}, @@ -19179,7 +19290,7 @@ async def delete_config_general_settings( proxy_config.settings.apply_db_row("general_settings", general_settings) if is_resource_list("general_settings", data.field_name): stored_endpoints: Final = general_settings.get("pass_through_endpoints") - await proxy_config._serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) + await proxy_config.serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ()) asyncio.create_task( create_config_audit_log( "general_settings", "deleted", before_general_settings, general_settings, user_api_key_dict @@ -19705,7 +19816,7 @@ async def get_model_cost_map_reload_status( Get the status of the scheduled model cost map reload job. """ # Read-only status check — admin viewers can read. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", @@ -19757,7 +19868,7 @@ async def get_model_cost_map_source( - model_count: number of models in the currently loaded cost map """ # Read-only source info — admin viewers can read. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", @@ -19979,7 +20090,7 @@ async def get_anthropic_beta_headers_reload_status( Get the status of the scheduled Anthropic beta headers reload job. """ # Read-only status — admin viewers can read. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", @@ -20080,7 +20191,7 @@ async def get_adaptive_router_state( which deployment it came from. """ # Read-only state — admin viewers can read. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={"error": CommonProxyErrors.not_allowed_access.value}, @@ -20247,7 +20358,7 @@ async def _stream_mcp_asgi_response(handle_fn, scope: dict, receive) -> "Streami Call an ASGI MCP handler and return a StreamingResponse so SSE/streaming works. asyncio.create_task copies the current context, so any ContextVar set before - this call (e.g. _mcp_active_toolset_id) is visible inside the handler task. + this call (e.g. mcp_active_toolset_id) is visible inside the handler task. """ from starlette.responses import StreamingResponse @@ -20387,8 +20498,8 @@ async def toolset_mcp_route(toolset_name: str, request: Request): global_mcp_server_manager, ) from litellm.proxy._experimental.mcp_server.server import ( - _mcp_active_toolset_id, handle_streamable_http_mcp, + mcp_active_toolset_id, ) if prisma_client is None: @@ -20404,11 +20515,11 @@ async def toolset_mcp_route(toolset_name: str, request: Request): scope: Final = dict(request.scope) scope["path"] = "/mcp" - token: Final = _mcp_active_toolset_id.set(toolset.toolset_id) + token: Final = mcp_active_toolset_id.set(toolset.toolset_id) try: return await _stream_mcp_asgi_response(handle_streamable_http_mcp, scope, request.receive) finally: - _mcp_active_toolset_id.reset(token) + mcp_active_toolset_id.reset(token) except HTTPException as e: raise e @@ -20494,7 +20605,7 @@ async def _is_mcp_access_group_cached(name: str) -> bool: cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key) if cached is not None: return bool(cached) - result: Final = bool(await MCPRequestHandler._get_mcp_servers_from_access_groups([name])) + result: Final = bool(await MCPRequestHandler.get_mcp_servers_from_access_groups([name])) await user_api_key_cache.async_set_cache( key=cache_key, value=result, @@ -20550,8 +20661,8 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request): # 3. Toolset name (cached) if prisma_client is not None: from litellm.proxy._experimental.mcp_server.server import ( - _mcp_active_toolset_id, handle_streamable_http_mcp, + mcp_active_toolset_id, ) toolset: Final = await global_mcp_server_manager.get_toolset_by_name_cached(prisma_client, mcp_server_name) @@ -20559,11 +20670,11 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request): scope: Final = dict(request.scope) scope["_original_path"] = scope.get("path", "") scope["path"] = "/mcp" - token: Final = _mcp_active_toolset_id.set(toolset.toolset_id) + token: Final = mcp_active_toolset_id.set(toolset.toolset_id) try: return await _stream_mcp_asgi_response(handle_streamable_http_mcp, scope, request.receive) finally: - _mcp_active_toolset_id.reset(token) + mcp_active_toolset_id.reset(token) # 4. MCP access group tag (cached) if await _is_mcp_access_group_cached(mcp_server_name): diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index da9a3033187..6b6ee7a9b67 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -221,10 +221,10 @@ def _load_endpoints() -> list[_EndpointEntry]: async def public_model_hub(): import litellm from litellm.proxy.health_endpoints._health_endpoints import ( - _convert_health_check_to_dict, + convert_health_check_to_dict, ) from litellm.proxy.proxy_server import ( - _get_model_group_info, + get_model_group_info, llm_router, prisma_client, ) @@ -234,7 +234,7 @@ async def public_model_hub(): model_groups: list[ModelGroupInfoProxy] = [] if litellm.public_model_groups is not None: - model_groups = _get_model_group_info( + model_groups = get_model_group_info( # rebind-ok: pre-existing rebinding on a rename-only line llm_router=llm_router, all_models_str=litellm.public_model_groups, model_group=None, @@ -248,7 +248,7 @@ async def public_model_hub(): for check in latest_checks: key = check.model_id if check.model_id else check.model_name if key: - health_check_dict = _convert_health_check_to_dict(check) + health_check_dict = convert_health_check_to_dict(check) health_checks_map[key] = health_check_dict if check.model_name: health_checks_map[check.model_name] = health_check_dict @@ -322,7 +322,7 @@ async def get_mcp_servers(): async def public_skill_hub(): """Return enabled (public) Claude Code skills — no auth required.""" from litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_marketplace import ( - _get_prisma_client, + get_prisma_client, ) from litellm.types.proxy.claude_code_endpoints import ( ListPluginsResponse, @@ -330,7 +330,7 @@ async def public_skill_hub(): ) try: - prisma_client: Final = await _get_prisma_client() + prisma_client: Final = await get_prisma_client() plugins: Final = await _plugin_table(prisma_client).find_many(where={"enabled": True}) items: Final = [] for plugin in plugins: @@ -366,7 +366,7 @@ async def public_skill_hub(): ) async def public_model_hub_info(): import litellm - from litellm.proxy.proxy_server import _title, version + from litellm.proxy.proxy_server import title, version try: from litellm_enterprise.proxy.proxy_server import EnterpriseProxyConfig @@ -376,7 +376,7 @@ async def public_model_hub_info(): custom_docs_description = None return PublicModelHubInfo( - docs_title=_title, + docs_title=title, custom_docs_description=custom_docs_description, litellm_version=version, useful_links=litellm.public_model_groups_links, diff --git a/litellm/proxy/public_endpoints/public_v1/model_hub.py b/litellm/proxy/public_endpoints/public_v1/model_hub.py index 93971a84547..2f5eeeb3d4b 100644 --- a/litellm/proxy/public_endpoints/public_v1/model_hub.py +++ b/litellm/proxy/public_endpoints/public_v1/model_hub.py @@ -186,7 +186,7 @@ MODEL_HUB_LIST_SPEC: Final[ListSpec[ModelGroupInfoProxy, ModelGroupInfoProxy]] = def _published_rows() -> Sequence[ModelGroupInfoProxy]: from litellm.proxy.proxy_server import ( - _get_model_group_info, # pyright: ignore[reportPrivateUsage] # /public/model_hub imports it the same way + get_model_group_info, # pyright: ignore[reportPrivateUsage] # /public/model_hub imports it the same way llm_router, ) @@ -202,7 +202,7 @@ def _published_rows() -> Sequence[ModelGroupInfoProxy]: if litellm.public_model_groups is None: return () return tuple( - _get_model_group_info( + get_model_group_info( llm_router=llm_router, all_models_str=litellm.public_model_groups, model_group=None, diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 0913d216678..a1a7d62899f 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -31,10 +31,12 @@ from litellm.proxy.common_request_processing import ( open_sse_before_first_byte, ttft_keepalive_interval, ) -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export get_form_data, + read_request_body, + safe_get_request_headers, ) from litellm.proxy.rag_endpoints.upload_security import ( MAX_UPLOAD_SIZE_BYTES, @@ -420,7 +422,7 @@ async def parse_rag_ingest_request( Returns: Tuple of (ingest_options, file_data, file_url, file_id) """ - headers: Final = _safe_get_request_headers(request) + headers: Final = safe_get_request_headers(request) content_type = headers.get("content-type", "") file_data: tuple[str, bytes, str] | None = None @@ -448,7 +450,7 @@ async def parse_rag_ingest_request( else: # JSON body - data: Final = await _read_request_body(request) + data: Final = await read_request_body(request) ingest_options = data.get("ingest_options", {}) file_url = data.get("file_url") file_id = data.get("file_id") @@ -770,7 +772,7 @@ async def rag_query( try: # Parse request body - data: Final = await _read_request_body(request) + data: Final = await read_request_body(request) # Extract required fields model: Final = data.get("model") diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index be2ac2ff33e..07b1abb62e4 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -16,7 +16,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_error_payload import ( error_status_code, openai_error_param, @@ -247,7 +250,7 @@ async def create_realtime_client_secret( data: dict = {} try: - body: Final = await _read_request_body(request=request) + body: Final = await read_request_body(request=request) req: Final = RealtimeClientSecretRequest(**body) model, session_data, session_type = await _prepare_client_secret_session( @@ -559,7 +562,7 @@ async def create_realtime_transcription_session( data: dict = {} try: - body: Final = await _read_request_body(request=request) + body: Final = await read_request_body(request=request) req: Final = RealtimeTranscriptionSessionRequest(**body) model: Final[str] = req.resolved_model() or "gpt-realtime-whisper" diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index eae621ba1ce..17ea2754821 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -30,9 +30,11 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth_websocket, ) from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing, create_response -from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_set_request_parsed_body, +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + _safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, + safe_set_request_parsed_body, ) from litellm.proxy.route_llm_request import raise_if_required_body_param_missing from litellm.types.llms.base import LiteLLMBaseModel @@ -184,12 +186,12 @@ async def _resolve_cursor_model_variant_before_auth(request: Request) -> None: from litellm.proxy.proxy_server import llm_router try: - raw_body: Final = await _read_request_body(request=request) + raw_body: Final = await read_request_body(request=request) except (json.JSONDecodeError, ProxyException): return resolved: Final = _resolve_cursor_model_variant(raw_body, llm_router) if resolved is not raw_body: - _safe_set_request_parsed_body(request=request, parsed_body=resolved) + safe_set_request_parsed_body(request=request, parsed_body=resolved) @router.post( @@ -244,7 +246,6 @@ async def responses_api( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, native_background_mode, @@ -252,6 +253,7 @@ async def responses_api( polling_via_cache_enabled, proxy_config, proxy_logging_obj, + read_request_body, redis_usage_cache, select_data_generator, user_api_base, @@ -263,7 +265,7 @@ async def responses_api( ) native_data_generator: Final = partial(select_data_generator, responses_stream_errors=True) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line # Check if polling via cache should be used for this request from litellm.proxy.response_polling.polling_handler import ( @@ -313,7 +315,7 @@ async def responses_api( ) raise_if_required_body_param_missing(route_type="aresponses", data=data, llm_router=llm_router) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -451,7 +453,7 @@ async def responses_api( return await create_response(generator=_blocked_stream(), media_type="text/event-stream", headers={}) return build_blocked_response(e) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -547,7 +549,7 @@ async def cursor_chat_completions( from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import ModelResponse - raw_body: Final = await _read_request_body(request=request) + raw_body: Final = await read_request_body(request=request) if _is_chat_completions_body(raw_body): # Genuine chat completions body (Cursor sends these for models whose BYOK it @@ -556,7 +558,7 @@ async def cursor_chat_completions( # empty messages stub alongside a real agent-mode input array normalized: Final = _normalize_tool_dialect(raw_body, to_chat=True) if normalized is not raw_body: - _safe_set_request_parsed_body(request=request, parsed_body=normalized) + safe_set_request_parsed_body(request=request, parsed_body=normalized) return await chat_completion( request=request, fastapi_response=fastapi_response, @@ -667,7 +669,7 @@ async def cursor_chat_completions( # Streaming responses are already transformed by cursor_select_data_generator return response except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -719,11 +721,11 @@ async def get_response( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, redis_usage_cache, select_data_generator, user_api_base, @@ -760,7 +762,7 @@ async def get_response( return state # Normal provider response flow - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["response_id"] = response_id processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: @@ -783,7 +785,7 @@ async def get_response( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -830,11 +832,11 @@ async def delete_response( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, redis_usage_cache, select_data_generator, user_api_base, @@ -869,7 +871,7 @@ async def delete_response( raise HTTPException(status_code=500, detail="Failed to delete polling response") # Normal provider response flow - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["response_id"] = response_id processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: @@ -892,7 +894,7 @@ async def delete_response( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -926,11 +928,11 @@ async def get_response_input_items( ): """List input items for a response.""" from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -940,7 +942,7 @@ async def get_response_input_items( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["response_id"] = response_id processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: @@ -963,7 +965,7 @@ async def get_response_input_items( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1009,11 +1011,11 @@ async def compact_response( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -1023,7 +1025,7 @@ async def compact_response( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: return await processor.base_process_llm_request( @@ -1045,7 +1047,7 @@ async def compact_response( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1162,7 +1164,7 @@ async def responses_input_tokens( Returns: `{"object": "response.input_tokens", "input_tokens": }` """ - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) model_name: Final = data.get("model") input_value: Final = data.get("input") if not isinstance(model_name, str) or not model_name: @@ -1240,11 +1242,11 @@ async def cancel_response( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, redis_usage_cache, select_data_generator, user_api_base, @@ -1283,7 +1285,7 @@ async def cancel_response( raise HTTPException(status_code=500, detail="Failed to cancel polling response") # Normal provider response flow - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["response_id"] = response_id processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: @@ -1306,7 +1308,7 @@ async def cancel_response( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1436,8 +1438,8 @@ async def _enforce_responses_ws_first_frame_model_auth( llm_router: "Router | None", ) -> None: from litellm.proxy.auth.user_api_key_auth import ( - _enforce_key_and_fallback_model_access, - _run_centralized_common_checks, + enforce_key_and_fallback_model_access, + run_centralized_common_checks, ) from litellm.proxy.proxy_server import ( general_settings, @@ -1456,7 +1458,7 @@ async def _enforce_responses_ws_first_frame_model_auth( return if user_custom_auth is not None and not general_settings.get("custom_auth_run_common_checks", False): return - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=user_api_key_dict, request_data=request_data, route=route, @@ -1464,7 +1466,7 @@ async def _enforce_responses_ws_first_frame_model_auth( llm_model_list=llm_model_list, llm_router=llm_router, ) - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=user_api_key_dict, request=request, request_data=request_data, @@ -1535,7 +1537,7 @@ async def responses_websocket_endpoint( "headers": headers_list, } request: Final = Request(scope=scope) - request._url = websocket.url + request._url = websocket.url # pyright: ignore[reportPrivateUsage] # Starlette WebSocket URL storage _body_bytes: Final = json.dumps({"model": resolved_model}).encode() diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index d6d673c6c76..30dfb5fb8a6 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -412,7 +412,9 @@ async def add_shared_session_to_data(data: dict) -> None: "SESSION REUSE: Shared aiohttp session is None after re-check, recreating..." ) try: - new_session = await proxy_server._initialize_shared_aiohttp_session() + new_session = ( # rebind-ok: pre-existing rebinding on a rename-only line + await proxy_server.initialize_shared_aiohttp_session() + ) except Exception: verbose_proxy_logger.exception("SESSION REUSE: Exception during shared session recreation") new_session = None diff --git a/litellm/proxy/search_endpoints/endpoints.py b/litellm/proxy/search_endpoints/endpoints.py index 1e37e3243ea..f5adc9be9f1 100644 --- a/litellm/proxy/search_endpoints/endpoints.py +++ b/litellm/proxy/search_endpoints/endpoints.py @@ -208,7 +208,7 @@ async def search( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/search_endpoints/search_tool_registry.py b/litellm/proxy/search_endpoints/search_tool_registry.py index fe119a9d44d..c8cd1100e08 100644 --- a/litellm/proxy/search_endpoints/search_tool_registry.py +++ b/litellm/proxy/search_endpoints/search_tool_registry.py @@ -12,10 +12,11 @@ from pydantic import TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy.auth.master_key_boot_check import SALT_KEY_ENV_VAR -from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _get_salt_key, +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports + _get_salt_key, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export decrypt_if_encrypted_with, encrypt_value_helper, + get_salt_key, ) from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.proxy.utils import PrismaClient @@ -85,7 +86,7 @@ def encrypt_search_tool_litellm_params(litellm_params: Mapping[str, object]) -> def _search_tool_plaintext(value: str) -> str | None: - signing_key: Final = _get_salt_key() + signing_key: Final = get_salt_key() return None if signing_key is None else decrypt_if_encrypted_with(value, signing_key) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index f9235a33b42..4c9df95cd8d 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -438,10 +438,10 @@ async def invalidate_budget_reservation_counters( if budget_reservation is None: return - from litellm.proxy.proxy_server import _invalidate_spend_counter + from litellm.proxy.proxy_server import invalidate_spend_counter for counter_key in get_reserved_counter_keys(budget_reservation=budget_reservation): - await _invalidate_spend_counter(counter_key=counter_key) + await invalidate_spend_counter(counter_key=counter_key) async def release_or_invalidate_budget_reservation( @@ -946,19 +946,19 @@ async def _reservation_counter_loaded(counter: _BudgetCounter, fail_closed_budge async def _initialize_reservation_counter(counter: _BudgetCounter) -> None: from litellm.proxy.proxy_server import ( - _ensure_spend_counter_initialized, - _ensure_window_spend_counter_initialized, - _invalidate_spend_counter, + ensure_spend_counter_initialized, + ensure_window_spend_counter_initialized, + invalidate_spend_counter, ) try: if counter.source_cache_key is not None: - await _ensure_spend_counter_initialized( + await ensure_spend_counter_initialized( counter_key=counter.counter_key, source_cache_key=counter.source_cache_key, ) elif counter.spend_log_entity_id is not None and counter.window_start is not None: - initialized: Final = await _ensure_window_spend_counter_initialized( + initialized: Final = await ensure_window_spend_counter_initialized( counter_key=counter.counter_key, entity_type=counter.entity_type, entity_id=counter.spend_log_entity_id, @@ -980,7 +980,7 @@ async def _initialize_reservation_counter(counter: _BudgetCounter) -> None: exc_info=True, ) try: - await _invalidate_spend_counter(counter_key=counter.counter_key) + await invalidate_spend_counter(counter_key=counter.counter_key) except Exception: verbose_proxy_logger.warning( "Failed to invalidate spend counter after budget reservation failure for %s", @@ -1062,7 +1062,7 @@ async def _reserve_counters( ) -> tuple[float | None, ...] | None: """One INCRBYFLOAT pipeline reserves every counter. When it fails each counter is dropped, and one that cannot be dropped is released instead in case its increment landed, so nothing is left to release by the caller.""" - from litellm.proxy.proxy_server import _invalidate_spend_counter, run_spend_counter_pipeline + from litellm.proxy.proxy_server import invalidate_spend_counter, run_spend_counter_pipeline if not counters: return () @@ -1081,7 +1081,7 @@ async def _reserve_counters( ) for counter, entry in zip(counters, entries): try: - await _invalidate_spend_counter(counter_key=counter.counter_key) + await invalidate_spend_counter(counter_key=counter.counter_key) except Exception: verbose_proxy_logger.warning( "Failed to invalidate spend counter after budget reservation failure for %s", @@ -1190,11 +1190,11 @@ async def _reseed_reserved_entry(item: _EntryAdjustment, actual_cost: float) -> reconcile: the optimistic delta no longer applies, so reseed from the DB floor and add the settled cost, since increment_spend_counters skips reserved keys. The reconcile runs before this request's spend is enqueued to the DB, so the reseeded floor excludes it.""" - from litellm.proxy.proxy_server import _increment_spend_counter_cache, reseed_spend_counter_from_db + from litellm.proxy.proxy_server import increment_spend_counter_cache, reseed_spend_counter_from_db reseeded: Final = await reseed_spend_counter_from_db(counter_key=item.counter_key) if reseeded and actual_cost > 0: - await _increment_spend_counter_cache(counter_key=item.counter_key, increment=actual_cost) + await increment_spend_counter_cache(counter_key=item.counter_key, increment=actual_cost) async def _counter_can_apply_adjustment( @@ -1230,9 +1230,9 @@ async def _release_applied_entries_best_effort( if counter_key is None: continue try: - from litellm.proxy.proxy_server import _invalidate_spend_counter + from litellm.proxy.proxy_server import invalidate_spend_counter - await _invalidate_spend_counter(counter_key=counter_key) + await invalidate_spend_counter(counter_key=counter_key) except Exception: verbose_proxy_logger.exception( "Failed to invalidate partial budget reservation counter during exception cleanup" diff --git a/litellm/proxy/spend_tracking/cloudzero_endpoints.py b/litellm/proxy/spend_tracking/cloudzero_endpoints.py index 6c0d11a2174..66a1934b40b 100644 --- a/litellm/proxy/spend_tracking/cloudzero_endpoints.py +++ b/litellm/proxy/spend_tracking/cloudzero_endpoints.py @@ -17,7 +17,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, +) from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.prisma_protocols import TableActions from litellm.types.proxy.cloudzero_endpoints import ( @@ -149,7 +152,7 @@ async def get_cloudzero_settings( Only admin users (Proxy Admin or Admin Viewer) can view CloudZero settings. """ # Validation — Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={"error": CommonProxyErrors.not_allowed_access.value}, diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index af33a3e528b..f40500028a2 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -449,7 +449,7 @@ async def spend_key_fn( "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" ) - if _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + if is_admin_view_safe(user_api_key_dict=user_api_key_dict): return await prisma_client.get_data(table_name="key", query_type="find_all") caller_user_id: Final = user_api_key_dict.user_id @@ -522,7 +522,7 @@ async def spend_user_fn( "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" ) - if not _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + if not is_admin_view_safe(user_api_key_dict=user_api_key_dict): caller_user_id: Final = user_api_key_dict.user_id if not caller_user_id: return [] @@ -1244,7 +1244,7 @@ async def get_spend_capture_rate( """ from litellm.proxy.proxy_server import prisma_client - if not _is_admin_view_safe(user_api_key_dict): + if not is_admin_view_safe(user_api_key_dict): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Only proxy admins can read the capture rate") if prisma_client is None: raise HTTPException( @@ -1863,7 +1863,7 @@ def _resolve_spend_report_scope( viewers) may request any scope. """ if requested: - if requested != caller_value and not _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + if requested != caller_value and not is_admin_view_safe(user_api_key_dict=user_api_key_dict): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=f"Not authorized to view spend for a {scope_name} other than your own", @@ -1887,7 +1887,7 @@ async def _resolve_org_spend_report_scope( Callable by proxy admins (any organization) and org admins of the target organization; every other caller is a 403 from ``_verify_org_access``. """ - from litellm.proxy.management_endpoints.organization_endpoints import _verify_org_access + from litellm.proxy.management_endpoints.organization_endpoints import verify_org_access target_org = organization_id or user_api_key_dict.org_id if target_org is None: @@ -1895,7 +1895,7 @@ async def _resolve_org_spend_report_scope( status_code=status.HTTP_400_BAD_REQUEST, detail="No organization_id associated with this API key; pass an organization_id query param", ) - await _verify_org_access( + await verify_org_access( organization_id=target_org, user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, @@ -2212,10 +2212,10 @@ async def global_view_spend_tags( ) -async def _get_spend_report_for_time_range( +async def get_spend_report_for_time_range( start_date: str, end_date: str, -): +) -> tuple[Sequence[_TeamSpendRow] | None, Sequence[_TagSpendRow] | None] | None: from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -2274,6 +2274,9 @@ async def _get_spend_report_for_time_range( verbose_proxy_logger.error("Exception in _get_daily_spend_reports %s", e) +_get_spend_report_for_time_range: Final = get_spend_report_for_time_range + + @router.post( "/spend/calculate", tags=["Budget & Spend Tracking"], @@ -2652,7 +2655,7 @@ async def ui_view_spend_logs( ) try: - is_admin_view: Final = _is_admin_view_safe(user_api_key_dict=user_api_key_dict) + is_admin_view: Final = is_admin_view_safe(user_api_key_dict=user_api_key_dict) is_request_id_lookup: Final = request_id is not None and not is_v2 is_search_lookup: Final = search is not None search_owns_window: Final = is_search_lookup and not is_v2 @@ -3387,7 +3390,7 @@ async def ui_view_request_response_for_request_id( """ from litellm.proxy.proxy_server import prisma_client - caller_is_admin: Final = _is_admin_view_safe(user_api_key_dict=user_api_key_dict) + caller_is_admin: Final = is_admin_view_safe(user_api_key_dict=user_api_key_dict) if not caller_is_admin: if prisma_client is None: raise HTTPException( @@ -4532,7 +4535,7 @@ async def ui_view_session_spend_logs( read_scope: Final = ( AllRows() - if _is_admin_view_safe(user_api_key_dict=user_api_key_dict) + if is_admin_view_safe(user_api_key_dict=user_api_key_dict) else await _spend_log_read_scope(user_api_key_dict, log_team_lookup) if _can_user_view_spend_log(user_api_key_dict=user_api_key_dict) else OwnedRows(user_api_key_dict.user_id) @@ -4793,7 +4796,7 @@ def _span_type_sql_condition(span_type: str | None) -> str | None: return _SPAN_TYPE_SQL_CONDITIONS.get(span_type) -def _is_admin_view_safe(user_api_key_dict: UserAPIKeyAuth) -> bool: +def is_admin_view_safe(user_api_key_dict: UserAPIKeyAuth) -> bool: """ Safely determine if the current user has admin view permissions. Defaults to False on any exception. @@ -4810,6 +4813,9 @@ def _is_admin_view_safe(user_api_key_dict: UserAPIKeyAuth) -> bool: return False +_is_admin_view_safe: Final = is_admin_view_safe + + async def _can_team_member_view_log( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index b71b834a31c..9d52f7d71bb 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -86,7 +86,7 @@ def _get_max_string_length_prompt_in_db() -> int: return DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB -def _is_master_key(api_key: str | None, _master_key: str | None) -> bool: +def is_master_key(api_key: str | None, _master_key: str | None) -> bool: """ Raw-only constant-time master-key comparison. The hashed form is never considered equivalent — only the raw master-key string matches. @@ -96,6 +96,9 @@ def _is_master_key(api_key: str | None, _master_key: str | None) -> bool: return secrets.compare_digest(api_key, _master_key) +_is_master_key: Final = is_master_key + + _HASHED_JWT_RE = re.compile(r"hashed-jwt-[a-fA-F0-9]{64}") _NON_SECRET_KEY_ALIASES: Final = frozenset( { @@ -1411,7 +1414,7 @@ def _redact_prompt_fields_in_guardrail_entry( return {**redacted, "guardrail_response": preserved_stats} -def _sanitize_error_information_for_spend_logs( +def sanitize_error_information_for_spend_logs( error_information: StandardLoggingPayloadErrorInformation | None, original_exception: BaseException | None = None, ) -> StandardLoggingPayloadErrorInformation | None: @@ -1452,6 +1455,9 @@ def _sanitize_error_information_for_spend_logs( return cast(StandardLoggingPayloadErrorInformation, sanitized) +_sanitize_error_information_for_spend_logs: Final = sanitize_error_information_for_spend_logs + + def _convert_to_json_serializable_dict(obj: object, visited: set[int] | None = None, max_depth: int = 20) -> object: """ Convert object to JSON-serializable dict, handling Pydantic models safely. diff --git a/litellm/proxy/spend_tracking/vantage_endpoints.py b/litellm/proxy/spend_tracking/vantage_endpoints.py index 8ea28f42cd9..d3858f95243 100644 --- a/litellm/proxy/spend_tracking/vantage_endpoints.py +++ b/litellm/proxy/spend_tracking/vantage_endpoints.py @@ -18,7 +18,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports + _user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + user_api_key_has_admin_view, +) from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.prisma_protocols import TableActions from litellm.types.proxy.vantage_endpoints import ( @@ -160,7 +163,7 @@ async def get_vantage_settings( Only admin users (Proxy Admin or Admin Viewer) can view Vantage settings. """ # Admin Viewer follows the read-parity rule. - if not _user_has_admin_view(user_api_key_dict): + if not user_api_key_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={"error": CommonProxyErrors.not_allowed_access.value}, diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 368c2e9758e..c410e66f0e3 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -1262,7 +1262,9 @@ async def update_sso_settings( if isinstance(stored, str): stored = json.loads(stored) if isinstance(stored, dict): - before_sso_data = proxy_config._decrypt_db_variables(stored) + before_sso_data = ( # rebind-ok: pre-existing rebinding on a rename-only line + proxy_config.decrypt_db_variables(stored) + ) # Load existing config config: Final = await proxy_config.get_config() @@ -1286,7 +1288,7 @@ async def update_sso_settings( # Clear environment variable if value is null/empty os.environ.pop(env_var_name, None) - encrypted_sso_data: Final = proxy_config._encrypt_env_variables(environment_variables=sso_data) + encrypted_sso_data: Final = proxy_config.encrypt_env_variables(environment_variables=sso_data) # Save to dedicated SSO table await _stored_sso_settings_db(SSOConfigRepository(prisma_client)).upsert( @@ -1554,7 +1556,7 @@ async def update_mcp_semantic_filter_settings( from litellm.proxy.proxy_server import prisma_client, proxy_config if prisma_client is not None: - await proxy_config._init_semantic_filter_settings_in_db(prisma_client=prisma_client) + await proxy_config.init_semantic_filter_settings_in_db(prisma_client=prisma_client) except Exception as e: verbose_proxy_logger.warning("Failed to reinitialize MCP semantic filter settings immediately: %s", e) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 3a3c3fa2928..3e73e2dee40 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -47,7 +47,7 @@ from typing import ( runtime_checkable, ) -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import Never, ReadOnly, TypedDict from litellm import _custom_logger_compatible_callbacks_literal from litellm.constants import ( @@ -190,7 +190,11 @@ from litellm.proxy.db.health_check_latest import ( fetch_latest_health_checks, fetch_latest_health_checks_for_models, ) -from litellm.proxy.db.log_db_metrics import _is_exception_related_to_db, log_db_metrics +from litellm.proxy.db.log_db_metrics import ( # noqa: F401, RUF100 # legacy module exports + _is_exception_related_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + is_exception_related_to_db, + log_db_metrics, +) from litellm.proxy.db.pgbouncer import database_url_is_pooled from litellm.proxy.db.prisma_client import ( PrismaWrapper, @@ -213,15 +217,21 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai resolve_endpoint_translation, ) from litellm.proxy.hooks import PROXY_HOOKS, get_proxy_hook -from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck -from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, +from litellm.proxy.hooks.cache_control_check import ( # noqa: F401, RUF100 # legacy module exports + PROXY_CacheControlCheck, + _PROXY_CacheControlCheck, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) -from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, +from litellm.proxy.hooks.parallel_request_limiter import ( # noqa: F401, RUF100 # legacy module exports + PROXY_MaxParallelRequestsHandler, + _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) -from litellm.proxy.hooks.sensitive_data_routing import ( - _PROXY_SensitiveDataRoutingHandler, +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( # noqa: F401, RUF100 # legacy module exports + PROXY_MaxParallelRequestsHandler_v3, + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export +) +from litellm.proxy.hooks.sensitive_data_routing import ( # noqa: F401, RUF100 # legacy module exports + PROXY_SensitiveDataRoutingHandler, + _PROXY_SensitiveDataRoutingHandler, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export ) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_guardrails_from_auth_metadata from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at @@ -1219,8 +1229,8 @@ class ProxyLogging: self.file_usage_cache: Final = InternalUsageCache( dual_cache=DualCache(in_memory_cache=InMemoryCache(max_size_in_memory=FILE_USAGE_MAX_TRACKED_COUNTERS)) ) - self.max_parallel_request_limiter = _PROXY_MaxParallelRequestsHandler(self.internal_usage_cache) - self.cache_control_check = _PROXY_CacheControlCheck() + self.max_parallel_request_limiter = PROXY_MaxParallelRequestsHandler(self.internal_usage_cache) + self.cache_control_check = PROXY_CacheControlCheck() self.alerting: list[str] | None = None self.alerting_threshold: float = 300 # default to 5 min. threshold self.alert_types: list[AlertType] = DEFAULT_ALERT_TYPES @@ -1554,6 +1564,13 @@ class ProxyLogging: ] return synthetic_data + def convert_mcp_to_llm_format( + self, + request_obj: MCPPreCallRequestObject, + kwargs: Mapping[str, object], + ) -> dict[str, object]: + return self._convert_mcp_to_llm_format(request_obj, kwargs) + def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> MCPPreCallResponseObject | None: """ Convert LLM guardrail result back to MCP response format. @@ -1742,7 +1759,7 @@ class ProxyLogging: } return result - def _create_mcp_request_object_from_kwargs(self, kwargs: dict) -> "MCPPreCallRequestObject": + def create_mcp_request_object_from_kwargs(self, kwargs: dict) -> "MCPPreCallRequestObject": """ Helper function to create MCPPreCallRequestObject from kwargs for standard pre_call_hook. """ @@ -1761,7 +1778,9 @@ class ProxyLogging: hidden_params=HiddenParams(), ) - def _convert_mcp_hook_response_to_kwargs(self, response_data: dict | None, original_kwargs: dict) -> dict: + _create_mcp_request_object_from_kwargs = create_mcp_request_object_from_kwargs + + def convert_mcp_hook_response_to_kwargs(self, response_data: dict | None, original_kwargs: dict) -> dict: """ Helper function to convert pre_call_hook response back to kwargs for MCP usage. @@ -1790,6 +1809,8 @@ class ProxyLogging: return modified_kwargs + _convert_mcp_hook_response_to_kwargs = convert_mcp_hook_response_to_kwargs + async def process_pre_call_hook_response(self, response, data, call_type): if isinstance(response, Exception): raise response @@ -2355,7 +2376,7 @@ class ProxyLogging: """ if request_metadata.get("_guardrail_pipelines"): return True - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() if caps.has_content_enforcer: return True probe: Final = {"metadata": dict(request_metadata)} @@ -2455,7 +2476,7 @@ class ProxyLogging: # otherwise makes deep copies return the original object. needs_raw_request_snapshot: Final = any( isinstance(cb, CustomGuardrail) and cb.scan_raw_request - for cb in ProxyLogging._callback_capabilities().resolved_callbacks + for cb in ProxyLogging.callback_capabilities().resolved_callbacks ) raw_request_snapshot: Final[dict | None] = ( # mutable-ok: same request-payload shape as data independent_snapshot(data) if needs_raw_request_snapshot else None @@ -2476,7 +2497,7 @@ class ProxyLogging: frozenset() if skip_guardrails else pipeline_managed_guardrail_names(data, "pre_call") ) - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() # Skip the per-request callback walk entirely when nothing in # ``litellm.callbacks`` overrides ``async_pre_call_hook`` and no # CustomGuardrail is configured. Saves the loop overhead + @@ -2544,7 +2565,7 @@ class ProxyLogging: call_type=call_type, endpoint_type=endpoint_type, ) - if isinstance(_callback, _PROXY_MaxParallelRequestsHandler_v3) + if isinstance(_callback, PROXY_MaxParallelRequestsHandler_v3) else await _callback.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=self.call_details["user_api_key_cache"], @@ -2694,7 +2715,7 @@ class ProxyLogging: if exc.sticky_session_routing: sensitive_routing_hook: Final = self.get_proxy_hook("sensitive_data_routing") - if isinstance(sensitive_routing_hook, _PROXY_SensitiveDataRoutingHandler): + if isinstance(sensitive_routing_hook, PROXY_SensitiveDataRoutingHandler): await sensitive_routing_hook.set_session_routing( session_id=exc.session_id, model=exc.route_to_model, @@ -2814,7 +2835,7 @@ class ProxyLogging: _callback_capabilities_cache: ClassVar[dict[tuple[int, tuple[int, ...]], "_CallbackCapabilities"]] = {} @staticmethod - def _callback_capabilities() -> "_CallbackCapabilities": + def callback_capabilities() -> "_CallbackCapabilities": """ Inspect ``litellm.callbacks`` once and answer the per-hook capability questions used to short-circuit no-op work on the chat-completions hot @@ -2916,6 +2937,8 @@ class ProxyLogging: cache[sig] = caps return caps + _callback_capabilities = callback_capabilities + @staticmethod def _stream_requires_guardrail_translation(user_api_key_dict: UserAPIKeyAuth) -> bool: route: Final = user_api_key_dict.request_route @@ -2928,18 +2951,18 @@ class ProxyLogging: @staticmethod def has_post_call_response_headers_callbacks() -> bool: - return ProxyLogging._callback_capabilities().has_post_call_response_headers + return ProxyLogging.callback_capabilities().has_post_call_response_headers @staticmethod def has_streaming_callbacks() -> bool: - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() return caps.has_iterator_override or caps.has_streaming_chunk_override or caps.has_guardrail @staticmethod def has_streaming_chunk_hook_overrides() -> bool: """True iff any callback overrides ``async_post_call_streaming_hook`` (the per-chunk hook, distinct from the iterator wrapper).""" - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() return caps.has_streaming_chunk_override or caps.has_guardrail def needs_iterator_wrap(self) -> bool: @@ -2947,19 +2970,19 @@ class ProxyLogging: through ``async_post_call_streaming_iterator_hook``. Instance method so tests can override the gate via ``MagicMock(spec=ProxyLogging)``. """ - return ProxyLogging._callback_capabilities().has_iterator_override + return ProxyLogging.callback_capabilities().has_iterator_override def needs_per_chunk_streaming_hook(self) -> bool: """Whether ``async_data_generator`` needs to call the per-chunk ``_apply_streaming_chunk_hooks`` for every emitted chunk. Instance method for the same reason as :py:meth:`needs_iterator_wrap`. """ - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() return caps.has_streaming_chunk_override or caps.has_guardrail @staticmethod def has_during_call_guardrails() -> bool: - return ProxyLogging._callback_capabilities().has_guardrail + return ProxyLogging.callback_capabilities().has_guardrail async def during_call_hook( self, @@ -2967,7 +2990,7 @@ class ProxyLogging: user_api_key_dict: UserAPIKeyAuth | None, call_type: CallTypesLiteral, ): - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() if not caps.has_guardrail and not caps.has_moderation_override: return data # Step 1: Collect all guardrail tasks to run in parallel @@ -3222,7 +3245,7 @@ class ProxyLogging: ) ) - logged_by_decorator: Final = call_type in _LOG_DB_METRICS_CALL_TYPES and _is_exception_related_to_db( + logged_by_decorator: Final = call_type in _LOG_DB_METRICS_CALL_TYPES and is_exception_related_to_db( original_exception ) if hasattr(self, "service_logging_obj") and not logged_by_decorator: @@ -3564,7 +3587,7 @@ class ProxyLogging: guardrail_callbacks, other_callbacks = _partition_post_call_callbacks() try: # Merge model-level guardrails before checking which guardrails to run - guardrail_data: Final = _check_and_merge_model_level_guardrails(data=data, llm_router=llm_router) + guardrail_data: Final = check_and_merge_model_level_guardrails(data=data, llm_router=llm_router) parallel_guardrails: Final[tuple[CustomGuardrail, ...]] = tuple( callback @@ -3724,7 +3747,7 @@ class ProxyLogging: (matching the inbound ``pre_mcp_call`` behavior) rather than being swallowed into an unguarded result. """ - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() if not caps.has_guardrail: return response @@ -3777,7 +3800,7 @@ class ProxyLogging: # cached detection makes the redundant interior guard cheap, but the # guard would still iterate every code path through this function so # keep it cheap and rely on the cached capability lookup. - if not ProxyLogging._callback_capabilities().has_post_call_response_headers: + if not ProxyLogging.callback_capabilities().has_post_call_response_headers: return merged_headers try: @@ -3819,7 +3842,7 @@ class ProxyLogging: async def hidden_by_listing_callbacks( self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str] ) -> frozenset[str]: - filters: Final = ProxyLogging._callback_capabilities().listed_models_filters + filters: Final = ProxyLogging.callback_capabilities().listed_models_filters if not filters: return frozenset() candidates: Final = tuple(model_names) @@ -3874,7 +3897,7 @@ class ProxyLogging: # active. ``get_response_string`` walks every choice/delta on the # chunk so paying it per chunk for no-op callbacks dominated stream # CPU time even after the iterator-chain fix. - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() if not caps.has_streaming_chunk_override and not caps.has_guardrail: return response @@ -3907,7 +3930,7 @@ class ProxyLogging: ## CHECK FOR MODEL-LEVEL GUARDRAILS (cached per-request) if not _guardrail_data_computed: - _cached_guardrail_data = _check_and_merge_model_level_guardrails( + _cached_guardrail_data = check_and_merge_model_level_guardrails( data=data, llm_router=llm_router ) _guardrail_data_computed = True @@ -3957,7 +3980,7 @@ class ProxyLogging: Covers: 1. /chat/completions """ - caps: Final = ProxyLogging._callback_capabilities() + caps: Final = ProxyLogging.callback_capabilities() post_call_pipelines: Final = _streamable_post_call_pipelines(request_data, user_api_key_dict) # Fast path: no real overrides. Internal proxy CustomLogger callbacks # (e.g. _PROXY_CacheControlCheck, ManagedFiles) inherit the default @@ -3972,15 +3995,15 @@ class ProxyLogging: raise except Exception as e: if not ProxyLogging._discard_deferred_stream_logging_for_failure(request_data, e): - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) raise - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) return from litellm.proxy.proxy_server import llm_router # Merge model-level guardrails before checking which guardrails to run - request_data = _check_and_merge_model_level_guardrails(data=request_data, llm_router=llm_router) + request_data = check_and_merge_model_level_guardrails(data=request_data, llm_router=llm_router) current_response = response stream_needs_translation: Final = ProxyLogging._stream_requires_guardrail_translation(user_api_key_dict) @@ -4056,7 +4079,7 @@ class ProxyLogging: except Exception as e: ProxyLogging._record_served_stream_output(request_data, served_chunks) if not ProxyLogging._discard_deferred_stream_logging_for_failure(request_data, e): - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) raise # Fire deferred logging AFTER all guardrail end-of-stream blocks @@ -4064,7 +4087,7 @@ class ProxyLogging: # its end-of-stream block (inside current_response), so by the time # we reach this point the metadata is fully populated. ProxyLogging._record_served_stream_output(request_data, served_chunks) - ProxyLogging._fire_deferred_stream_logging(request_data) + ProxyLogging.fire_deferred_stream_logging(request_data) async def _pipeline_gated_stream( self, @@ -4149,7 +4172,7 @@ class ProxyLogging: record_served_output_texts(logging_obj.model_call_details, served_stream_output_texts(served_chunks)) @staticmethod - def _fire_deferred_stream_logging(request_data: dict) -> None: + def fire_deferred_stream_logging(request_data: dict) -> None: """ Fire the deferred streaming logging callback after the full streaming pipeline (including guardrail end-of-stream blocks) has completed. @@ -4170,6 +4193,8 @@ class ProxyLogging: logging_obj._deferred_stream_complete_args = None asyncio.create_task(_deferred_cb(*_args)) + _fire_deferred_stream_logging = fire_deferred_stream_logging + @staticmethod def _discard_deferred_stream_logging_for_failure(request_data: Mapping[str, object], error: Exception) -> bool: """Drop the parked success dispatch for an assembled chat stream that ends in an error @@ -4187,7 +4212,7 @@ class ProxyLogging: logging_obj.record_assembled_response_for_failure(assembled) return True - async def _arelease_max_parallel_requests_on_disconnect( + async def arelease_max_parallel_requests_on_disconnect( self, user_api_key_dict: UserAPIKeyAuth, ) -> None: @@ -4207,17 +4232,19 @@ class ProxyLogging: double-decrement under the limiter's in-memory fallback. """ limiter: Final = self.get_proxy_hook("parallel_request_limiter") - if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + if not isinstance(limiter, PROXY_MaxParallelRequestsHandler_v3): return await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict) + _arelease_max_parallel_requests_on_disconnect = arelease_max_parallel_requests_on_disconnect + async def enforce_mcp_server_rate_limits( self, user_api_key_dict: UserAPIKeyAuth | None, server: "MCPServer", ) -> None: limiter: Final = self.get_proxy_hook("parallel_request_limiter") - if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + if not isinstance(limiter, PROXY_MaxParallelRequestsHandler_v3): return await limiter.enforce_mcp_server_rate_limits(user_api_key_dict, server) @@ -4964,7 +4991,9 @@ class PrismaClient: # check if plain text or hash if token is not None: if isinstance(token, str): - hashed_token = _hash_token_if_needed(token=token) + hashed_token = hash_token_if_needed( # rebind-ok: pre-existing rebinding on a rename-only line + token=token + ) verbose_proxy_logger.debug("PrismaClient: find_unique for token: %s", hashed_token) if query_type == "find_unique" and hashed_token is not None: if token is None: @@ -5026,7 +5055,7 @@ class PrismaClient: if token is not None: where_filter["token"] = {} if isinstance(token, str): - token = _hash_token_if_needed(token=token) + token = hash_token_if_needed(token=token) where_filter["token"]["in"] = [token] elif isinstance(token, list): hashed_tokens: Final[list[str]] = [] @@ -5196,7 +5225,9 @@ class PrismaClient: # check if plain text or hash if token is not None: if isinstance(token, str): - hashed_token = _hash_token_if_needed(token=token) + hashed_token = hash_token_if_needed( # rebind-ok: pre-existing rebinding on a rename-only line + token=token + ) verbose_proxy_logger.debug("PrismaClient: find_unique for token: %s", hashed_token) if query_type == "find_unique": if token is None: @@ -5529,7 +5560,7 @@ class PrismaClient: if token is not None: print_verbose(f"token: [set={token is not None}]") # check if plain text or hash - token = _hash_token_if_needed(token=token) + token = hash_token_if_needed(token=token) db_data["token"] = token include_object_permission: Final[LiteLLM_VerificationTokenInclude] = {"object_permission": True} response: Final = await VerificationTokenRepository(self).table.update( @@ -5877,7 +5908,7 @@ class PrismaClient: prisma_obj: Final = self.writer_db._original_prisma if prisma_obj.is_connected() is not True: return 0 - engine: Final = prisma_obj._engine + engine: Final = prisma_obj._engine # pyright: ignore[reportPrivateUsage] # Prisma engine internals process: Final = getattr(engine, "process", None) if engine is not None else None if process is not None: pid: Final[object] = process.pid @@ -7071,7 +7102,7 @@ class PrismaClient: ### HELPER FUNCTIONS ### -async def _cache_user_row(user_id: str, cache: DualCache, db: PrismaClient): +async def cache_user_row(user_id: str, cache: DualCache, db: PrismaClient) -> None: """ Check if a user_id exists in cache, if not retrieve it. @@ -7087,6 +7118,9 @@ async def _cache_user_row(user_id: str, cache: DualCache, db: PrismaClient): cache.set_cache(key=cache_key, value=cache_value, ttl=600) # store for 10 minutes +_cache_user_row: Final = cache_user_row + + def _should_use_smtp_ssl(smtp_port: int) -> bool: """ Port 465 expects an immediate TLS handshake (implicit SSL), so a plain @@ -7263,7 +7297,7 @@ async def migrate_passwords_to_scrypt_async(prisma_client) -> str: return f"Migrated {len(plaintext_users)} plaintext passwords to scrypt" -def _hash_token_if_needed(token: str) -> str: +def hash_token_if_needed(token: str) -> str: """ Hash the token if it's a string and starts with "sk-" @@ -7275,6 +7309,9 @@ def _hash_token_if_needed(token: str) -> str: return token +_hash_token_if_needed: Final = hash_token_if_needed + + async def enqueue_spend_logs( prisma_client: PrismaClient, logs: Sequence[Mapping[str, object]], @@ -7381,7 +7418,7 @@ class ProxyUpdateSpend: break except Exception as e: - await DBSpendUpdateWriter._handle_spend_update_failure( + await DBSpendUpdateWriter.handle_spend_update_failure( e=e, attempt=i, n_retry_times=n_retry_times, @@ -7476,7 +7513,7 @@ class ProxyUpdateSpend: raise await asyncio.sleep(1 << i) except Exception as e: - _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) + raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) finally: # Clean up logs_to_process only if we popped it (caller-owned otherwise) if popped_batch: @@ -7635,14 +7672,14 @@ async def update_daily_tag_spend( """ n_retry_times: Final = 3 try: - if proxy_logging_obj.db_spend_update_writer.redis_update_buffer._should_commit_spend_updates_to_redis(): - await proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis( + if proxy_logging_obj.db_spend_update_writer.redis_update_buffer.should_commit_spend_updates_to_redis(): + await proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis( prisma_client=prisma_client, n_retry_times=n_retry_times, proxy_logging_obj=proxy_logging_obj, ) else: - await proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db( + await proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db( prisma_client=prisma_client, n_retry_times=n_retry_times, proxy_logging_obj=proxy_logging_obj, @@ -7879,11 +7916,11 @@ async def _park_remaining_spend_logs(prisma_client: PrismaClient, proxy_logging_ ) -async def _monitor_spend_logs_queue( +async def monitor_spend_logs_queue( prisma_client: PrismaClient, db_writer_client: AsyncHTTPHandler | None, proxy_logging_obj: ProxyLogging, -): +) -> Never: """ Background task that monitors the spend_log_transactions queue size and triggers processing when the threshold is reached. @@ -7955,6 +7992,9 @@ async def _monitor_spend_logs_queue( await asyncio.sleep(current_interval) +_monitor_spend_logs_queue: Final = monitor_spend_logs_queue + + MAX_SPEND_LOG_ISOLATION_FAILURES_PER_BATCH: Final = 256 @@ -8028,7 +8068,7 @@ async def _create_spend_logs_with_poison_isolation( return await _create_spend_logs_with_poison_isolation(repo, rows[mid:], remaining) -def _raise_failed_update_spend_exception(e: Exception, start_time: float, proxy_logging_obj: ProxyLogging): +def raise_failed_update_spend_exception(e: Exception, start_time: float, proxy_logging_obj: ProxyLogging) -> Never: """ Raise an exception for failed update spend logs @@ -8052,13 +8092,16 @@ def _raise_failed_update_spend_exception(e: Exception, start_time: float, proxy_ raise e +_raise_failed_update_spend_exception: Final = raise_failed_update_spend_exception + + def _get_month_end_date(today: date) -> date: if today.month == 12: return date(today.year + 1, 1, 1) - timedelta(days=1) return date(today.year, today.month + 1, 1) - timedelta(days=1) -def _is_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None): +def is_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> bool: if soft_budget_limit is None: # If there's no limit, we can't exceed it. return False @@ -8085,7 +8128,10 @@ def _is_projected_spend_over_limit(current_spend: float, soft_budget_limit: floa return False -def _get_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> tuple | None: +_is_projected_spend_over_limit: Final = is_projected_spend_over_limit + + +def get_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> tuple | None: if soft_budget_limit is None: return None @@ -8117,7 +8163,10 @@ def _get_projected_spend_over_limit(current_spend: float, soft_budget_limit: flo return None -def _is_valid_team_configs(team_id=None, team_config=None, request_data=None): +_get_projected_spend_over_limit: Final = get_projected_spend_over_limit + + +def is_valid_team_configs(team_id=None, team_config=None, request_data=None) -> None: if team_id is None or team_config is None or request_data is None: return # check if valid model called for team @@ -8131,11 +8180,14 @@ def _is_valid_team_configs(team_id=None, team_config=None, request_data=None): return +_is_valid_team_configs: Final = is_valid_team_configs + + def _to_ns(dt): return int(dt.timestamp() * 1e9) -def _check_and_merge_model_level_guardrails( +def check_and_merge_model_level_guardrails( data: dict, llm_router: Router | None, trust_client_model_info: bool = True, @@ -8222,6 +8274,9 @@ def _check_and_merge_model_level_guardrails( return _merge_guardrails_with_existing(data, model_level_guardrails) +_check_and_merge_model_level_guardrails: Final = check_and_merge_model_level_guardrails + + def _merge_guardrails_with_existing(data: dict, model_level_guardrails: object) -> dict: """ Merge model-level guardrails with any existing guardrails in the request data. @@ -8270,7 +8325,7 @@ def get_error_message_str(e: Exception) -> str: return error_message -def _get_redoc_url() -> str | None: +def get_redoc_url() -> str | None: """ Get the Redoc URL from the environment variables. @@ -8287,7 +8342,10 @@ def _get_redoc_url() -> str | None: return "/redoc" -def _get_docs_url() -> str | None: +_get_redoc_url: Final = get_redoc_url + + +def get_docs_url() -> str | None: """ Get the docs (Swagger UI) URL from the environment variables. @@ -8304,7 +8362,10 @@ def _get_docs_url() -> str | None: return "/" -def _get_openapi_url() -> str | None: +_get_docs_url: Final = get_docs_url + + +def get_openapi_url() -> str | None: """ Get the OpenAPI JSON URL from the environment variables. @@ -8321,6 +8382,9 @@ def _get_openapi_url() -> str | None: return "/openapi.json" +_get_openapi_url: Final = get_openapi_url + + def _recreate_writer_on_read_only_transaction(prisma_client: "PrismaClient | None") -> None: if prisma_client is None: return @@ -8382,7 +8446,9 @@ def require_enterprise_license(feature: str | None = None) -> None: ) -_premium_user_check: Final = require_enterprise_license +premium_user_check: Final = require_enterprise_license + +_premium_user_check: Final = premium_user_check def is_known_model(model: str | None, llm_router: Router | None) -> bool: @@ -8611,11 +8677,11 @@ async def _get_access_group_models( proxy_logging_obj: Optional["ProxyLogging"], ) -> tuple[str, ...]: from litellm.proxy.auth.auth_checks import ( - _get_models_from_access_groups, get_authorized_resources_from_key_access_groups, + get_models_from_access_groups, ) - team_group_models: Final = await _get_models_from_access_groups( + team_group_models: Final = await get_models_from_access_groups( access_group_ids=(team_object.access_group_ids or ()) if team_object is not None else (), prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index c2203bf7f67..22721dc26bc 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -118,11 +118,11 @@ async def vector_store_search( https://platform.openai.com/docs/api-reference/vector-stores/search """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -132,7 +132,7 @@ async def vector_store_search( version, ) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line reject_caller_embedding_selection_params(payload=data, source="the search request body") data["vector_store_id"] = vector_store_id @@ -168,7 +168,7 @@ async def vector_store_search( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -198,11 +198,11 @@ async def vector_store_create( ``` """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -212,7 +212,7 @@ async def vector_store_create( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) # Check for target_model_names parameter target_model_names: Final = data.pop("target_model_names", None) @@ -275,7 +275,7 @@ async def vector_store_create( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -338,7 +338,7 @@ async def vector_store_retrieve( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -408,7 +408,7 @@ async def vector_store_list( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -431,11 +431,11 @@ async def vector_store_update( https://platform.openai.com/docs/api-reference/vector-stores/modify """ from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -445,7 +445,7 @@ async def vector_store_update( version, ) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line if "vector_store_id" not in data: data["vector_store_id"] = vector_store_id @@ -474,7 +474,7 @@ async def vector_store_update( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -537,7 +537,7 @@ async def vector_store_delete( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index e44aa022334..96eda99425d 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -564,11 +564,11 @@ async def vector_store_file_create( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -578,7 +578,7 @@ async def vector_store_file_create( version, ) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line data["vector_store_id"] = vector_store_id managed_vector_store: Final = await assert_user_can_access_vector_store_id( vector_store_id=vector_store_id, @@ -644,7 +644,7 @@ async def vector_store_file_create( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -751,7 +751,7 @@ async def vector_store_file_list( user_api_key_dict=user_api_key_dict, ) except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -858,7 +858,7 @@ async def vector_store_file_retrieve( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -968,7 +968,7 @@ async def vector_store_file_content( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -996,11 +996,11 @@ async def vector_store_file_update( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): from litellm.proxy.proxy_server import ( - _read_request_body, general_settings, llm_router, proxy_config, proxy_logging_obj, + read_request_body, select_data_generator, user_api_base, user_max_tokens, @@ -1010,7 +1010,7 @@ async def vector_store_file_update( version, ) - data = await _read_request_body(request=request) + data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line data["vector_store_id"] = vector_store_id data["file_id"] = file_id managed_vector_store: Final = await assert_user_can_access_vector_store_id( @@ -1075,7 +1075,7 @@ async def vector_store_file_update( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -1182,7 +1182,7 @@ async def vector_store_file_delete( return response except Exception as e: # noqa: BLE001 - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py index 523b669280e..724fb128a9a 100644 --- a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py +++ b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py @@ -21,8 +21,14 @@ import litellm from litellm.litellm_core_utils.url_utils import SSRFError, validate_url from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers -from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + safe_get_request_headers, +) +from litellm.proxy.litellm_pre_call_utils import ( # noqa: F401 # legacy module exports + _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + get_dynamic_logging_metadata, +) from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( create_pass_through_route, ) @@ -36,7 +42,7 @@ def create_request_copy(request: Request): return { "method": request.method, "url": str(request.url), - "headers": _safe_get_request_headers(request).copy(), + "headers": safe_get_request_headers(request).copy(), "cookies": request.cookies, "query_params": dict(request.query_params), } @@ -175,7 +181,7 @@ async def langfuse_proxy_route( user_api_key_dict: Final = await user_api_key_auth(request=request, api_key=f"Bearer {api_key}") - callback_settings_obj: Final[TeamCallbackMetadata | None] = _get_dynamic_logging_metadata( + callback_settings_obj: Final[TeamCallbackMetadata | None] = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index 9175d067920..2813333f113 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -9,7 +9,10 @@ from starlette.datastructures import UploadFile as StarletteUploadFile from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports + _read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export + read_request_body, +) from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_body, get_custom_llm_provider_from_request_headers, @@ -80,7 +83,7 @@ async def video_generation( ) # Read request body - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) if input_reference is not None: input_reference_file: Final = await batch_to_bytesio([input_reference]) if input_reference_file: @@ -108,7 +111,7 @@ async def video_generation( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -195,7 +198,7 @@ async def video_list( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -295,7 +298,7 @@ async def video_status( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -402,7 +405,7 @@ async def video_content( headers={"Content-Disposition": f"attachment; filename=video_{video_id}.mp4"}, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -458,7 +461,7 @@ async def video_remix( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["video_id"] = video_id decoded: Final = decode_video_id_with_provider(video_id) @@ -503,7 +506,7 @@ async def video_remix( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -560,7 +563,7 @@ async def video_create_character( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) video_file: Final = await batch_to_bytesio([video]) if video_file: data["video"] = video_file[0] @@ -608,7 +611,7 @@ async def video_create_character( ) return response except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -714,7 +717,7 @@ async def video_get_character( ) return response except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -767,7 +770,7 @@ async def video_edit( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) uploaded_video: Final = data.pop("video", None) if isinstance(uploaded_video, StarletteUploadFile): video_files: Final = await batch_to_bytesio((uploaded_video,)) @@ -816,7 +819,7 @@ async def video_edit( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, @@ -871,7 +874,7 @@ async def video_extension( version, ) - data: Final = await _read_request_body(request=request) + data: Final = await read_request_body(request=request) data["video_id"] = video_reference_to_id(data.pop("video", None)) decoded: Final = decode_video_id_with_provider(data["video_id"]) @@ -913,7 +916,7 @@ async def video_extension( version=version, ) except Exception as e: - raise await processor._handle_llm_api_exception( + raise await processor.handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 3293851fd9f..d3e3509b8f0 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -298,10 +298,10 @@ class LiteLLM_Proxy_MCP_Handler: granted_toolset_ids, ) from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_view, + user_api_key_has_admin_view, ) - if not _user_has_admin_view(user_api_key_auth) and toolset.toolset_id not in ( + if not user_api_key_has_admin_view(user_api_key_auth) and toolset.toolset_id not in ( await (granted_toolsets or granted_toolset_ids)(user_api_key_auth) ): verbose_logger.debug("Key does not have access to toolset '%s', skipping.", name) @@ -764,7 +764,7 @@ class LiteLLM_Proxy_MCP_Handler: mcp_server = global_mcp_server_manager.get_mcp_server_by_name( server_name - ) or global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name) + ) or global_mcp_server_manager.get_mcp_server_from_tool_name(tool_name) resolved_tool_name = ( _resolve_display_name_to_original(tool_name, [mcp_server]) if mcp_server else tool_name ) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index d486db72304..d2c240f8768 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -376,9 +376,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if raw_headers_from_request: headers_obj: Final = Headers(raw_headers_from_request) - self.mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers_obj) - self.mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj) - self.oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers_obj) + self.mcp_auth_header = MCPRequestHandler.get_mcp_auth_header_from_headers(headers_obj) + self.mcp_server_auth_headers = MCPRequestHandler.get_mcp_server_auth_headers_from_headers(headers_obj) + self.oauth2_headers = MCPRequestHandler.get_oauth2_headers_from_headers(headers_obj) # Also check if headers are provided in tools array (from request body) tools: Final[Sequence[object] | None] = self.original_request_params.get("tools") @@ -389,7 +389,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if tool_headers and isinstance(tool_headers, dict): # Merge tool headers into mcp_server_auth_headers headers_obj_from_tool = Headers(tool_headers) - tool_mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers( + tool_mcp_server_auth_headers = MCPRequestHandler.get_mcp_server_auth_headers_from_headers( headers_obj_from_tool ) diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index f4705e7132e..b16fe72b281 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -1121,9 +1121,15 @@ class ResponsesAPIRequestUtils: if raw_headers_from_request: headers_obj: Final = Headers(raw_headers_from_request) - mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers_obj) - mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers_obj) - oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers_obj) + mcp_auth_header = ( # rebind-ok: pre-existing rebinding on a rename-only line + MCPRequestHandler.get_mcp_auth_header_from_headers(headers_obj) + ) + mcp_server_auth_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line + MCPRequestHandler.get_mcp_server_auth_headers_from_headers(headers_obj) + ) + oauth2_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line + MCPRequestHandler.get_oauth2_headers_from_headers(headers_obj) + ) if tools: for tool in tools: @@ -1133,7 +1139,7 @@ class ResponsesAPIRequestUtils: # Merge tool headers into mcp_server_auth_headers # Extract server-specific headers from tool.headers headers_obj_from_tool = Headers(tool_headers) - tool_mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers( + tool_mcp_server_auth_headers = MCPRequestHandler.get_mcp_server_auth_headers_from_headers( headers_obj_from_tool ) if tool_mcp_server_auth_headers: diff --git a/tests/batches_tests/test_batch_rate_limits.py b/tests/batches_tests/test_batch_rate_limits.py index 06beed6b195..87d35e11c74 100644 --- a/tests/batches_tests/test_batch_rate_limits.py +++ b/tests/batches_tests/test_batch_rate_limits.py @@ -15,19 +15,19 @@ from litellm import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache -def _build_batch_limiter() -> _PROXY_BatchRateLimiter: +def _build_batch_limiter() -> PROXY_BatchRateLimiter: internal_usage_cache = InternalUsageCache(dual_cache=DualCache()) - return _PROXY_BatchRateLimiter( + return PROXY_BatchRateLimiter( internal_usage_cache=internal_usage_cache, - parallel_request_limiter=_PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter=PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ), ) @@ -136,7 +136,7 @@ async def test_batch_rate_limit_single_file(tmp_path): # Setup: Create internal usage cache and rate limiter dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) @@ -198,7 +198,7 @@ async def test_batch_rate_limit_single_file(tmp_path): # Reset cache for clean test dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() @@ -278,7 +278,7 @@ async def test_batch_rate_limit_multiple_requests(tmp_path): # Setup: Create internal usage cache and rate limiter dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) @@ -431,7 +431,7 @@ async def test_batch_rate_limiter_with_managed_files(tmp_path): # Setup: Create internal usage cache and rate limiter dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) @@ -662,7 +662,7 @@ async def test_batch_rate_limiter_managed_files_regression(): # Setup: Create batch rate limiter dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() @@ -685,10 +685,10 @@ async def test_batch_rate_limiter_managed_files_regression(): # Test 1: Verify managed file detection print("\n1. Verifying managed file detection...") from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, ) - is_managed = _is_base64_encoded_unified_file_id(managed_file_id) + is_managed = is_base64_encoded_unified_file_id(managed_file_id) assert is_managed, "Managed file should be detected correctly" print(" ✓ Managed file detected") diff --git a/tests/code_coverage_tests/unbounded_in_baseline.txt b/tests/code_coverage_tests/unbounded_in_baseline.txt index 129cf1b7797..77c7c362f39 100644 --- a/tests/code_coverage_tests/unbounded_in_baseline.txt +++ b/tests/code_coverage_tests/unbounded_in_baseline.txt @@ -38,7 +38,7 @@ litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval pr litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval prisma user_id.in `list(data.user_ids)` 0 litellm/proxy/management_endpoints/budget_management_endpoints.py info_budget prisma budget_id.in `data.budgets` 0 litellm/proxy/management_endpoints/common_utils.py _team_admin_can_invite_user prisma team_id.in `admin_user_obj.teams` 0 -litellm/proxy/management_endpoints/common_utils.py _user_has_admin_privileges prisma team_id.in `user_obj.teams` 0 +litellm/proxy/management_endpoints/common_utils.py user_has_admin_privileges prisma team_id.in `user_obj.teams` 0 litellm/proxy/management_endpoints/customer_endpoints.py delete_end_user prisma user_id.in `data.user_ids` 0 litellm/proxy/management_endpoints/customer_endpoints.py delete_end_user prisma user_id.in `data.user_ids` 1 litellm/proxy/management_endpoints/internal_user_endpoints.py _check_user_info_v2_access prisma team_id.in `caller_user.teams` 0 @@ -61,7 +61,7 @@ litellm/proxy/management_endpoints/key_management_endpoints.py _build_key_filter litellm/proxy/management_endpoints/key_management_endpoints.py _build_key_filter_conditions prisma team_id.in `member_only_team_ids` 0 litellm/proxy/management_endpoints/key_management_endpoints.py _build_key_filter_conditions prisma team_id.in `member_team_ids` 0 litellm/proxy/management_endpoints/key_management_endpoints.py _fetch_user_team_objects prisma team_id.in `complete_user_info.teams` 0 -litellm/proxy/management_endpoints/key_management_endpoints.py _list_key_helper prisma user_id.in `all_ids` 0 +litellm/proxy/management_endpoints/key_management_endpoints.py list_key_helper prisma user_id.in `all_ids` 0 litellm/proxy/management_endpoints/key_management_endpoints.py bulk_update_team_keys prisma token.in `hashed_key_ids` 0 litellm/proxy/management_endpoints/key_management_endpoints.py delete_key_aliases prisma key_alias.in `key_aliases` 0 litellm/proxy/management_endpoints/key_management_endpoints.py delete_verification_tokens prisma token.in `hashed_tokens` 0 diff --git a/tests/guardrails_tests/test_presidio_pii.py b/tests/guardrails_tests/test_presidio_pii.py index d8a1885b833..595cc95fa9f 100644 --- a/tests/guardrails_tests/test_presidio_pii.py +++ b/tests/guardrails_tests/test_presidio_pii.py @@ -5,7 +5,7 @@ from unittest.mock import patch import litellm from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, PresidioPerRequestConfig, ) from litellm.types.guardrails import PiiEntityType, PiiAction @@ -24,7 +24,7 @@ async def test_presidio_with_blocked_entities(): PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, # This entity should be masked } - presidio_guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail = OPTIONAL_PresidioPIIMasking( pii_entities_config=pii_entities_config, presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"), presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"), @@ -64,7 +64,7 @@ async def test_presidio_pre_call_hook_with_blocked_entities(): PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, # This entity should be masked } - presidio_guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail = OPTIONAL_PresidioPIIMasking( pii_entities_config=pii_entities_config, presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"), presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"), @@ -195,10 +195,10 @@ async def test_presidio_pii_masking_logging_output_only_logged_response_guardrai assert len(litellm.guardrail_name_config_map) == 1 - pii_masking_obj: Optional[_OPTIONAL_PresidioPIIMasking] = None + pii_masking_obj: Optional[OPTIONAL_PresidioPIIMasking] = None for callback in litellm.callbacks: print(f"CALLBACK: {callback}") - if isinstance(callback, _OPTIONAL_PresidioPIIMasking): + if isinstance(callback, OPTIONAL_PresidioPIIMasking): pii_masking_obj = callback assert pii_masking_obj is not None diff --git a/tests/integration/spend/test_redis_ttl_preserving_token_increment.py b/tests/integration/spend/test_redis_ttl_preserving_token_increment.py index 3f623280c17..8b7a21f73d5 100644 --- a/tests/integration/spend/test_redis_ttl_preserving_token_increment.py +++ b/tests/integration/spend/test_redis_ttl_preserving_token_increment.py @@ -8,7 +8,7 @@ from redis import Redis from litellm.caching.caching import DualCache from litellm.caching.redis_cache import RedisCache from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache from litellm.types.caching import RedisPipelineIncrementOperation diff --git a/tests/local_testing/test_caching.py b/tests/local_testing/test_caching.py index 0c68720473d..b8c7a4c189b 100644 --- a/tests/local_testing/test_caching.py +++ b/tests/local_testing/test_caching.py @@ -1744,7 +1744,7 @@ async def test_redis_proxy_batch_redis_get_cache(): from litellm.caching.caching import Cache, DualCache from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.hooks.batch_redis_get import _PROXY_BatchRedisRequests + from litellm.proxy.hooks.batch_redis_get import PROXY_BatchRedisRequests litellm.cache = Cache( type="redis", @@ -1755,7 +1755,7 @@ async def test_redis_proxy_batch_redis_get_cache(): ) batch_redis_get_obj = ( - _PROXY_BatchRedisRequests() + PROXY_BatchRedisRequests() ) # overrides the .async_get_cache method user_api_key_cache = DualCache() diff --git a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py index f5b758843be..7a3109839d8 100644 --- a/tests/pass_through_unit_tests/test_pass_through_unit_tests.py +++ b/tests/pass_through_unit_tests/test_pass_through_unit_tests.py @@ -308,7 +308,7 @@ async def test_pass_through_request_logging_failure_with_stream( # Patch both the logging handler and the httpx client with ( patch( - "litellm.proxy.pass_through_endpoints.streaming_handler.PassThroughStreamingHandler._route_streaming_logging_to_handler", + "litellm.proxy.pass_through_endpoints.streaming_handler.PassThroughStreamingHandler.route_streaming_logging_to_handler", new=mock_logging_failure, ), patch( diff --git a/tests/test_litellm_rust/cache/test_python_cache.py b/tests/test_litellm_rust/cache/test_python_cache.py index 3a2e3b8c143..490643191ea 100644 --- a/tests/test_litellm_rust/cache/test_python_cache.py +++ b/tests/test_litellm_rust/cache/test_python_cache.py @@ -12,10 +12,10 @@ from litellm.caching.caching_handler import ( ) from litellm.proxy._types import Litellm_EntityType, UserAPIKeyAuth from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, model_budget_spend_cache_key, ) -from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 +from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import InternalUsageCache from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.rust_bridge import runtime @@ -57,8 +57,8 @@ async def test_cache_hit_keeps_model_budget_spend_but_accounts_for_usage( monkeypatch.setenv("LITELLM_RUST", "1" if native else "0") litellm.cache = Cache() if legacy else _v2.Cache.memory() counters: Final = litellm.DualCache() - budget: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(counters) - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3( + budget: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(counters) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3( InternalUsageCache(counters), model_group_resolver=lambda model: model ) recorder: Final = RecordingLogger() @@ -133,8 +133,8 @@ async def test_response_cache_backend_does_not_control_coordination( ) recording_server.expected_requests = 2 if backend == "disabled" else 1 counters: Final = litellm.DualCache() - budget: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(counters) - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3( + budget: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(counters) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3( InternalUsageCache(counters), model_group_resolver=lambda model: model ) key_hash: Final = "b" * 64 diff --git a/tests/test_presidio_latency.py b/tests/test_presidio_latency.py index 40a2cc42b25..bb676ca051e 100644 --- a/tests/test_presidio_latency.py +++ b/tests/test_presidio_latency.py @@ -3,7 +3,7 @@ import aiohttp import pytest from unittest.mock import MagicMock, patch from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) @@ -14,7 +14,7 @@ async def test_sanity_presidio_session_reuse_main_thread(): Verify that Presidio guardrail reuses sessions in the main thread. This ensures we don't break existing session pooling functionality. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_analyzer_api_base="http://mock-analyzer", presidio_anonymizer_api_base="http://mock-anonymizer", @@ -28,9 +28,7 @@ async def test_sanity_presidio_session_reuse_main_thread(): session_creations += 1 original_init(self, *args, **kwargs) - with patch.object( - aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True - ): + with patch.object(aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True): for _ in range(10): async with presidio._get_session_iterator() as session: pass @@ -51,7 +49,7 @@ async def test_bug_presidio_session_explosion_background_thread_causes_latency() """ import threading - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_analyzer_api_base="http://mock-analyzer", presidio_anonymizer_api_base="http://mock-anonymizer", @@ -68,9 +66,7 @@ async def test_bug_presidio_session_explosion_background_thread_causes_latency() session_creations += 1 original_init(self, *args, **kwargs) - with patch.object( - aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True - ): + with patch.object(aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True): for _ in range(10): async with presidio._get_session_iterator() as session: pass diff --git a/tests/unit/batches/test_batch_utils.py b/tests/unit/batches/test_batch_utils.py index bd8a546c0ce..293c38e9f20 100644 --- a/tests/unit/batches/test_batch_utils.py +++ b/tests/unit/batches/test_batch_utils.py @@ -1184,7 +1184,7 @@ async def test_output_file_content_fetches_and_parses(monkeypatch): return type("R", (), {"content": b'{"a": 1}\n{"b": 2}'})() monkeypatch.setattr(files_main, "afile_content", fake_afile_content) - monkeypatch.setattr(cu, "_is_base64_encoded_unified_file_id", lambda fid: False) + monkeypatch.setattr(cu, "is_base64_encoded_unified_file_id", lambda fid: False) result = await bu._fetch_batch_output_file_content( _batch("file-out"), @@ -1217,7 +1217,7 @@ async def test_output_file_content_unified_file_id_extraction(monkeypatch): monkeypatch.setattr(files_main, "afile_content", fake_afile_content) monkeypatch.setattr( cu, - "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", lambda fid: "litellm_proxy;llm_output_file_id,real-file-99;rest", ) diff --git a/tests/unit/caching/test_request_redis_batch_post_call.py b/tests/unit/caching/test_request_redis_batch_post_call.py index 8c8c9df5926..5de785e44db 100644 --- a/tests/unit/caching/test_request_redis_batch_post_call.py +++ b/tests/unit/caching/test_request_redis_batch_post_call.py @@ -34,7 +34,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( TOKEN_INCREMENT_SCRIPT, ParallelSlotAcquisition, RequestRateLimiterStash, - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.spend_tracking.spend_counter_batch import PendingSpendIncrement from litellm.proxy.utils import InternalUsageCache @@ -99,10 +99,10 @@ def _names(client: FakeClient, index: int = 0) -> list[str]: return [command[0] for command in client.pipelines[index].commands] -def _limiter(redis_cache: FakeRedisCache) -> _PROXY_MaxParallelRequestsHandler_v3: +def _limiter(redis_cache: FakeRedisCache) -> PROXY_MaxParallelRequestsHandler_v3: dual_cache = DualCache() dual_cache.attach_redis_cache(redis_cache) - return _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache=dual_cache)) + return PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache=dual_cache)) def _slot_stash(slot_id: str, *counter_keys: str) -> RequestRateLimiterStash: 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 81a2a1718e9..69ef67e595b 100644 --- a/tests/unit/caching/test_request_redis_batch_pre_call.py +++ b/tests/unit/caching/test_request_redis_batch_pre_call.py @@ -18,14 +18,14 @@ from litellm._internal_context import current_service_target from litellm.caching.dual_cache import DualCache from litellm.caching.redis_batch import active_request_redis_batches, request_redis_batch_scope from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable -from litellm.proxy.auth.auth_checks import _cache_team_object +from litellm.proxy.auth.auth_checks import cache_team_object from litellm.proxy.auth.auth_object_prefetch import _CacheEntry, _write_back, prefetch_identity_keys from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.hooks.parallel_request_limiter_v3 import ( CHECK_AND_INCREMENT_BY_N_SCRIPT, RateLimitDescriptor, RateLimitUnverifiableError, - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache from litellm.router_utils.cooldown_cache import ROUTER_COOLDOWNS_TARGET, CooldownCache @@ -46,9 +46,9 @@ def sha_of(script: str) -> str: return hashlib.sha1(script.encode()).hexdigest() # noqa: S324 -def _limiter(redis_cache: FakeRedisCache, fail_closed: bool = False) -> _PROXY_MaxParallelRequestsHandler_v3: +def _limiter(redis_cache: FakeRedisCache, fail_closed: bool = False) -> PROXY_MaxParallelRequestsHandler_v3: dual_cache = DualCache() - limiter = _PROXY_MaxParallelRequestsHandler_v3( + limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(dual_cache=dual_cache), fail_closed_resolver=lambda: fail_closed, ) @@ -64,7 +64,7 @@ def _descriptor(key: str, value: str, rpm: int) -> RateLimitDescriptor: return {"key": key, "value": value, "rate_limit": {"requests_per_unit": rpm}} -def _refunds(limiter: _PROXY_MaxParallelRequestsHandler_v3) -> list[tuple[str, float]]: +def _refunds(limiter: PROXY_MaxParallelRequestsHandler_v3) -> list[tuple[str, float]]: refund_script = limiter.window_guarded_token_increment_script assert isinstance(refund_script, AsyncMock) return [(call.kwargs["keys"][1], call.kwargs["args"][1]) for call in refund_script.await_args_list] @@ -1029,7 +1029,7 @@ async def test_a_team_refresh_inside_a_request_sends_its_set_and_alias_del_in_on proxy_logging_obj.internal_usage_cache = InternalUsageCache(dual_cache=usage_cache) team = LiteLLM_TeamTableCachedObj(team_id="t1", team_alias="alpha") with request_redis_batch_scope() as request: - await _cache_team_object("t1", team, cache, proxy_logging_obj) + await cache_team_object("t1", team, cache, proxy_logging_obj) assert [c[:2] for c in client.pipelines[0].commands] == [("SET", "team_id:t1"), ("DEL", "team_alias:alpha")], ( "the alias DEL must reach Redis before the refresh returns, or another request can refill memory from it" ) diff --git a/tests/unit/containers/test_container_proxy_ownership.py b/tests/unit/containers/test_container_proxy_ownership.py index dd21faa1059..780ae9bf6d1 100644 --- a/tests/unit/containers/test_container_proxy_ownership.py +++ b/tests/unit/containers/test_container_proxy_ownership.py @@ -425,7 +425,7 @@ async def test_should_validate_owner_and_forward_decoded_id_for_multipart_upload async def base_process_llm_request(self, **kwargs): return captured["data"] - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise kwargs["e"] monkeypatch.setattr( @@ -499,7 +499,7 @@ async def test_should_forward_decoded_container_id_for_proxy_retrieve(monkeypatc async def base_process_llm_request(self, **kwargs): return captured["data"] - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise kwargs["e"] monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor) @@ -554,7 +554,7 @@ async def test_should_record_container_owner_inside_create_endpoint(monkeypatch) async def base_process_llm_request(self, **kwargs): return response - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise kwargs["e"] record_owner = AsyncMock(return_value=response) @@ -608,7 +608,7 @@ async def test_should_not_route_owner_record_errors_through_llm_error_handler( async def base_process_llm_request(self, **kwargs): return _container("cntr_provider") - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise AssertionError("ownership errors should not use LLM error handler") monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor) @@ -668,7 +668,7 @@ async def test_should_return_response_when_owner_recording_raises_unexpected( async def base_process_llm_request(self, **kwargs): return created - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise AssertionError("upstream-create errors only") monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor) @@ -772,7 +772,7 @@ async def test_should_forward_decoded_container_id_for_proxy_delete(monkeypatch) async def base_process_llm_request(self, **kwargs): return captured["data"] - async def _handle_llm_api_exception(self, **kwargs): + async def handle_llm_api_exception(self, **kwargs): raise kwargs["e"] monkeypatch.setattr(endpoints, "ProxyBaseLLMRequestProcessing", FakeProcessor) diff --git a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index 78545d3fb62..de6a5601d5c 100644 --- a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -1711,7 +1711,7 @@ async def test_initialize_remaining_budget_metrics_exception_handling( "litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams" ) as mock_get_teams, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper" ) as mock_list_keys, ): # Make get_paginated_teams raise an exception @@ -1786,7 +1786,7 @@ async def test_initialize_api_key_budget_metrics(prometheus_logger): with ( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper" ) as mock_list_keys, ): # Create mock key data with proper datetime objects for budget_reset_at diff --git a/tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py b/tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py index d21bb046549..a93e1368c49 100644 --- a/tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/enterprise/proxy/auth/test_user_api_key_auth.py @@ -22,9 +22,7 @@ async def test_enterprise_custom_auth_mode_on(): mock_user_auth = AsyncMock(return_value={"user_id": "test-user"}) request = MagicMock(spec=Request) - with patch( - "litellm_enterprise.proxy.proxy_server.custom_auth_settings", {"mode": "on"} - ): + with patch("litellm_enterprise.proxy.proxy_server.custom_auth_settings", {"mode": "on"}): result = await enterprise_custom_auth(request, "test-api-key", mock_user_auth) assert result == {"user_id": "test-user"} mock_user_auth.assert_called_once_with(request, "test-api-key") @@ -36,9 +34,7 @@ async def test_enterprise_custom_auth_mode_auto_with_error(): mock_user_auth = AsyncMock(side_effect=Exception("Auth failed")) request = MagicMock(spec=Request) - with patch( - "litellm_enterprise.proxy.proxy_server.custom_auth_settings", {"mode": "auto"} - ): + with patch("litellm_enterprise.proxy.proxy_server.custom_auth_settings", {"mode": "auto"}): result = await enterprise_custom_auth(request, "test-api-key", mock_user_auth) assert result is None mock_user_auth.assert_called_once_with(request, "test-api-key") @@ -61,9 +57,7 @@ async def test_enterprise_custom_auth_returns_string(): patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), ): # Verify the key is correctly handled in _user_api_key_auth_builder - with patch( - "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key" - ) as mock_get_key_object: + with patch("litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key") as mock_get_key_object: mock_get_key_object.return_value = UserAPIKeyAuth( token="sk-test-key", user_role="internal_user", @@ -72,10 +66,10 @@ async def test_enterprise_custom_auth_returns_string(): ) # Call _user_api_key_auth_builder with the returned key - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder try: - auth_obj = await _user_api_key_auth_builder( + auth_obj = await user_api_key_auth_builder( request=request, api_key="my-custom-key", azure_api_key_header="", diff --git a/tests/unit/enterprise/proxy/hooks/test_managed_files.py b/tests/unit/enterprise/proxy/hooks/test_managed_files.py index 7d8c8b5c425..11d9dbec49c 100644 --- a/tests/unit/enterprise/proxy/hooks/test_managed_files.py +++ b/tests/unit/enterprise/proxy/hooks/test_managed_files.py @@ -11,7 +11,7 @@ from litellm.caching import DualCache from litellm.proxy._types import CallTypes from litellm.proxy.openai_files_endpoints.common_utils import ( BATCH_CREATE_HIDDEN_PARAM, - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, encode_file_id_with_model, ) @@ -365,7 +365,7 @@ async def test_async_post_call_success_hook_for_unified_finetuning_job(): ) assert isinstance(response, LiteLLMFineTuningJob) - assert _is_base64_encoded_unified_file_id(response.id) + assert is_base64_encoded_unified_file_id(response.id) @pytest.mark.asyncio @@ -603,7 +603,7 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi response=batch, ) - decoded_output_file_id = _is_base64_encoded_unified_file_id( + decoded_output_file_id = is_base64_encoded_unified_file_id( cast(LiteLLMBatch, response).output_file_id ) assert decoded_output_file_id @@ -691,7 +691,7 @@ async def test_error_file_id_for_failed_batch(): assert cast(LiteLLMBatch, response).error_file_id is not None assert not cast(LiteLLMBatch, response).error_file_id.startswith("error-") # Verify it's a base64 encoded managed file ID - assert _is_base64_encoded_unified_file_id( + assert is_base64_encoded_unified_file_id( cast(LiteLLMBatch, response).error_file_id ) @@ -754,7 +754,7 @@ async def test_async_post_call_success_hook_twice_assert_no_unique_violation(): assert task.exception() is None, f"Error: {task.exception()}" assert isinstance(response, LiteLLMBatch) - assert _is_base64_encoded_unified_file_id(response.id) + assert is_base64_encoded_unified_file_id(response.id) # second retrieve batch tasks = [] @@ -2762,7 +2762,7 @@ async def test_return_unified_file_id_includes_expires_at(): assert result.filename == "test.jsonl" assert result.bytes == 1234 assert result.created_at == 1234567890 - assert _is_base64_encoded_unified_file_id(result.id) + assert is_base64_encoded_unified_file_id(result.id) # ============================================================================ diff --git a/tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py b/tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py index 6e9c3c0354b..93ed7420b4c 100644 --- a/tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py +++ b/tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py @@ -13,7 +13,7 @@ import pytest from unittest.mock import AsyncMock, MagicMock, patch from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, ) @@ -32,7 +32,7 @@ async def test_should_resolve_raw_input_file_id_to_unified(): contains a record for that raw ID, the retrieve endpoint should resolve it to the unified file ID. """ - unified_batch_id = _is_base64_encoded_unified_file_id(B64_UNIFIED_BATCH_ID) + unified_batch_id = is_base64_encoded_unified_file_id(B64_UNIFIED_BATCH_ID) assert unified_batch_id, "Test setup: batch_id should decode as unified" from litellm.types.utils import LiteLLMBatch diff --git a/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py b/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py index 34103449dad..ebeb514ffff 100644 --- a/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py +++ b/tests/unit/integrations/open_telemetry/test_otel_admin_endpoints.py @@ -105,9 +105,7 @@ def test_key_generate_failure_stamps_server_span( ) -def test_key_generate_success_stamps_server_span( - server_span_factory, otel_with_exporter -): +def test_key_generate_success_stamps_server_span(server_span_factory, otel_with_exporter): otel, exporter = otel_with_exporter server_span = server_span_factory(KEY_GENERATE_PATH) @@ -230,9 +228,7 @@ def test_management_wrapper_success_ends_server_span_without_http_request( ) -def test_management_wrapper_failure_ends_server_span( - server_span_factory, otel_with_exporter, monkeypatch -): +def test_management_wrapper_failure_ends_server_span(server_span_factory, otel_with_exporter, monkeypatch): """When the handler raises, the wrapper must route through the failure hook and stamp + end the parent SERVER span with the error status — even for an ``http_request``-less handler (route falls back to ``func.__name__``).""" @@ -249,9 +245,7 @@ def test_management_wrapper_failure_ends_server_span( raise HttpStatusException(500, "boom") with pytest.raises(HttpStatusException): - asyncio.run( - failing_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span)) - ) + asyncio.run(failing_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))) assert_server_span_attrs( exporter, @@ -261,9 +255,7 @@ def test_management_wrapper_failure_ends_server_span( ) -def test_management_wrapper_success_with_http_request( - server_span_factory, otel_with_exporter, monkeypatch -): +def test_management_wrapper_success_with_http_request(server_span_factory, otel_with_exporter, monkeypatch): """Cover the branch where the handler DOES declare ``http_request``: the route comes from ``http_request.url.path`` and the body is read from it.""" import litellm.proxy.proxy_server as proxy_server @@ -276,7 +268,7 @@ def test_management_wrapper_success_with_http_request( async def _fake_body(request=None): return {"team_alias": "t"} - monkeypatch.setattr(mgmt_utils, "_read_request_body", _fake_body) + monkeypatch.setattr(mgmt_utils, "read_request_body", _fake_body) server_span = server_span_factory("/team/new") http_request = MagicMock() @@ -302,9 +294,7 @@ def test_management_wrapper_success_with_http_request( ) -def test_management_wrapper_noop_when_otel_logger_absent( - server_span_factory, otel_with_exporter, monkeypatch -): +def test_management_wrapper_noop_when_otel_logger_absent(server_span_factory, otel_with_exporter, monkeypatch): """When no OTEL logger is registered, the helper early-returns and no SERVER span is exported — and the handler result is still returned unchanged.""" import litellm.proxy.proxy_server as proxy_server @@ -320,17 +310,13 @@ def test_management_wrapper_noop_when_otel_logger_absent( async def fake_fn(data=None, user_api_key_dict=None): return {"ok": True} - result = asyncio.run( - fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span)) - ) + result = asyncio.run(fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))) assert result == {"ok": True} assert get_server_span(exporter) is None -def test_management_wrapper_swallows_post_success_errors( - server_span_factory, otel_with_exporter, monkeypatch -): +def test_management_wrapper_swallows_post_success_errors(server_span_factory, otel_with_exporter, monkeypatch): """A failure in post-success bookkeeping (cache invalidation, alerting) must not propagate — the handler result is returned regardless (non-blocking).""" import litellm.proxy.proxy_server as proxy_server @@ -351,8 +337,6 @@ def test_management_wrapper_swallows_post_success_errors( async def fake_fn(data=None, user_api_key_dict=None): return {"ok": True} - result = asyncio.run( - fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span)) - ) + result = asyncio.run(fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))) assert result == {"ok": True} diff --git a/tests/unit/integrations/test_shadow_eval_logger.py b/tests/unit/integrations/test_shadow_eval_logger.py index e3f059a7941..55796f7bbcf 100644 --- a/tests/unit/integrations/test_shadow_eval_logger.py +++ b/tests/unit/integrations/test_shadow_eval_logger.py @@ -1522,7 +1522,7 @@ class TestShadowPipeline: monkeypatch.setattr( auth_checks, - "_virtual_key_max_budget_check", + "virtual_key_max_budget_check", AsyncMock(side_effect=BudgetExceededError(current_cost=11.0, max_budget=10.0)), ) router = _router() 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 9a686ba9ab3..ffe50c137d2 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -840,12 +840,12 @@ async def test_ahealth_check_without_mode_reports_the_real_failure( def test_update_litellm_params_for_health_check(): """ - Test if _update_litellm_params_for_health_check correctly: + Test if update_litellm_params_for_health_check correctly: 1. Updates messages with a random message 2. Updates model name when health_check_model is provided 3. Updates voice when health_check_voice is provided for audio_speech mode """ - from litellm.proxy.health_check import _update_litellm_params_for_health_check + from litellm.proxy.health_check import update_litellm_params_for_health_check model_info = {"health_check_model": "gpt-5-mini"} litellm_params = { @@ -853,7 +853,7 @@ def test_update_litellm_params_for_health_check(): "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "messages" in updated_params assert isinstance(updated_params["messages"], list) @@ -865,7 +865,7 @@ def test_update_litellm_params_for_health_check(): "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "messages" in updated_params assert isinstance(updated_params["messages"], list) @@ -876,7 +876,7 @@ def test_update_litellm_params_for_health_check(): "model": "gpt-5.5", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "voice" in updated_params assert updated_params["voice"] == "en-US-JennyNeural" @@ -885,7 +885,7 @@ def test_update_litellm_params_for_health_check(): "model": "gpt-5.5", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "voice" in updated_params assert updated_params["voice"] == "alloy" @@ -894,7 +894,7 @@ def test_update_litellm_params_for_health_check(): "model": "gpt-5.5", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "voice" not in updated_params model_info = {} @@ -902,28 +902,28 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "anthropic.claude-sonnet-4-5-20250929-v1:0" litellm_params = { "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "us.anthropic.claude-haiku-4-5-20251001-v1:0" litellm_params = { "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "us.anthropic.claude-haiku-4-5-20251001-v1:0" litellm_params = { "model": "openai/gpt-5.5", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "openai/gpt-5.5" cris_prefixes = ["us.", "eu.", "apac.", "jp.", "au.", "us-gov.", "global."] @@ -932,7 +932,7 @@ def test_update_litellm_params_for_health_check(): "model": f"bedrock/{prefix}anthropic.claude-3-haiku-20240307-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check( + updated_params = update_litellm_params_for_health_check( model_info, litellm_params ) assert ( @@ -943,21 +943,21 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/us-east-2/us.anthropic.claude-3-haiku-20240307-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "us.anthropic.claude-3-haiku-20240307-v1:0" litellm_params = { "model": "bedrock/us-gov-east-1/anthropic.claude-instant-v1", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "anthropic.claude-instant-v1" litellm_params = { "model": "bedrock/llama/arn:aws:bedrock:us-east-1:123:imported-model/abc", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "llama/arn:aws:bedrock:us-east-1:123:imported-model/abc" @@ -967,7 +967,7 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/deepseek_r1/arn:aws:bedrock:us-west-2:456:imported-model/xyz", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "deepseek_r1/arn:aws:bedrock:us-west-2:456:imported-model/xyz" @@ -977,7 +977,7 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/converse/us.anthropic.claude-haiku-4-5-20251001-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "converse/us.anthropic.claude-haiku-4-5-20251001-v1:0" @@ -987,14 +987,14 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/invoke/us-west-2/anthropic.claude-instant-v1", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["model"] == "invoke/anthropic.claude-instant-v1" litellm_params = { "model": "bedrock/arn:aws:bedrock:eu-central-1:000:application-inference-profile/abc", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "arn:aws:bedrock:eu-central-1:000:application-inference-profile/abc" @@ -1004,7 +1004,7 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/us-west-2/llama/arn:aws:bedrock:us-east-1:123:imported-model/abc", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "llama/arn:aws:bedrock:us-east-1:123:imported-model/abc" @@ -1014,7 +1014,7 @@ def test_update_litellm_params_for_health_check(): "model": "bedrock/converse/us-west-2/eu.anthropic.claude-3-sonnet-20240229-v1:0", "api_key": "fake_key", } - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert ( updated_params["model"] == "converse/eu.anthropic.claude-3-sonnet-20240229-v1:0" ) diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index 3fbc425da7f..caa8708ee9a 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -28,6 +28,7 @@ import litellm from litellm._internal_context import in_post_response_phase from litellm._logging import session_id_var, trace_id_var, verbose_logger from litellm._service_logger import ServiceLogging +from litellm.caching.caching import DualCache from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE, REDACTED_BY_LITELLM, SENTRY_PII_DENYLIST from litellm.cost_calculator import ocr_batch_cost from litellm.integrations.custom_logger import CustomLogger @@ -42,8 +43,10 @@ from litellm.litellm_core_utils.litellm_logging import ( from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck -from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler +from litellm.proxy.hooks.cache_control_check import PROXY_CacheControlCheck +from litellm.proxy.hooks.max_iterations_limiter import PROXY_MaxIterationsHandler +from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 +from litellm.proxy.utils import InternalUsageCache from litellm.types.llms.openai import ResponseAPIUsage, ResponseCompletedEvent, ResponsesAPIResponse from litellm.types.utils import ( CallTypes, @@ -11074,6 +11077,24 @@ def test_litellm_logging_no_log_param(monkeypatch, disable_no_log_param): else: assert should_run is False + proxy_callback = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) + should_run_proxy_callback = litellm_logging_obj.should_run_callback( + callback=proxy_callback, + litellm_params={"no-log": True}, + event_hook="success_handler", + ) + assert should_run_proxy_callback is True + + from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles + + managed_files_callback = PROXY_LiteLLMManagedFiles(DualCache(), prisma_client=MagicMock()) + should_run_managed_files_callback = litellm_logging_obj.should_run_callback( + callback=managed_files_callback, + litellm_params={"no-log": True}, + event_hook="success_handler", + ) + assert should_run_managed_files_callback is True + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") def test_get_callback_name(): @@ -11108,7 +11129,7 @@ def test_is_internal_litellm_proxy_callback(): """ logging = setup_logging() - assert logging._is_internal_litellm_proxy_callback(_PROXY_MaxIterationsHandler) == True + assert logging._is_internal_litellm_proxy_callback(PROXY_MaxIterationsHandler) == True # Test non-internal callbacks def regular_callback(): @@ -11141,7 +11162,7 @@ def test_should_run_sync_callbacks_for_async_calls(): assert logging._should_run_sync_callbacks_for_async_calls() == True # Test with internal callback only - litellm.success_callback = [_PROXY_MaxIterationsHandler] + litellm.success_callback = [PROXY_MaxIterationsHandler] assert logging._should_run_sync_callbacks_for_async_calls() == False @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @@ -11153,8 +11174,8 @@ def test_remove_internal_litellm_callbacks(): callbacks = [ regular_callback, - _PROXY_MaxIterationsHandler, - _PROXY_CacheControlCheck, + PROXY_MaxIterationsHandler, + PROXY_CacheControlCheck, "string_callback", ] @@ -11162,8 +11183,8 @@ def test_remove_internal_litellm_callbacks(): assert len(filtered) == 2 # Should only keep regular_callback and string_callback assert regular_callback in filtered assert "string_callback" in filtered - assert _PROXY_MaxIterationsHandler not in filtered - assert _PROXY_CacheControlCheck not in filtered + assert PROXY_MaxIterationsHandler not in filtered + assert PROXY_CacheControlCheck not in filtered @pytest.mark.asyncio async def test_background_interaction_completion_logs_while_in_progress_handler_is_parked(monkeypatch): diff --git a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 85e4c055f07..9c0fb5c6a98 100644 --- a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -182,7 +182,7 @@ class TestAnthropicMessagesHandlerStreamingRequestData: with ( patch.object(handler, "_check_streaming_has_ended", return_value=True), patch( - "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler.build_complete_streaming_response", return_value=mock_response, ), ): @@ -241,7 +241,7 @@ class TestAnthropicMessagesHandlerStreamingOutputProcessing: with ( patch.object(handler, "_check_streaming_has_ended", return_value=True), patch( - "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler.build_complete_streaming_response", return_value=None, ), ): @@ -1353,7 +1353,7 @@ class TestAnthropicMessagesHandlerInputProcessing: with ( patch.object(handler, "_check_streaming_has_ended", return_value=True), patch( - "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler.build_complete_streaming_response", return_value=mock_response, ), ): @@ -1400,7 +1400,7 @@ class TestAnthropicMessagesHandlerInputProcessing: with ( patch.object(handler, "_check_streaming_has_ended", return_value=True), patch( - "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler.build_complete_streaming_response", return_value=mock_response, ), ): diff --git a/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py b/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py index a94c47d73f1..ced53794cc2 100644 --- a/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py +++ b/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py @@ -1651,10 +1651,10 @@ async def test_summary_model_denied_when_user_over_model_budget(): import inspect from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, ) - real_params = inspect.signature(_PROXY_VirtualKeyModelMaxBudgetLimiter.is_user_within_model_budget).parameters + real_params = inspect.signature(PROXY_VirtualKeyModelMaxBudgetLimiter.is_user_within_model_budget).parameters for kwarg in ("user_id", "user_model_max_budget", "model"): assert kwarg in real_params, f"compact.py passes {kwarg}=, which the limiter no longer accepts" @@ -1767,7 +1767,7 @@ class _FakeRateLimiter: self._raises = raises self.read_only_checked = False - def _create_rate_limit_descriptors(self, **kwargs): + def create_rate_limit_descriptors(self, **kwargs): return [ { "key": "api_key", @@ -1921,12 +1921,12 @@ async def test_summary_model_allowed_while_the_caller_holds_the_keys_only_parall caller's own in-flight slot must not trip a ``max_parallel_requests`` gauge.""" from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import InternalUsageCache, hash_token messages = _simple_messages() mock_call = AsyncMock(return_value=_make_mock_response("ok")) - limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) + limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) auth = UserAPIKeyAuth( api_key=hash_token("sk-compact-parallel-slot"), max_parallel_requests=1, models=["all-proxy-models"] ) @@ -2061,11 +2061,11 @@ async def test_summary_model_denied_when_team_over_model_budget(): import inspect from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, ) real_params = inspect.signature( - _PROXY_VirtualKeyModelMaxBudgetLimiter.is_team_within_model_budget + PROXY_VirtualKeyModelMaxBudgetLimiter.is_team_within_model_budget ).parameters for kwarg in ("team_id", "team_model_max_budget", "key_model_max_budget", "model"): assert kwarg in real_params, f"compact.py passes {kwarg}=, which the limiter does not accept" diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py b/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py index 8690d2ee68f..b88a57f440d 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py @@ -289,7 +289,7 @@ async def test_cached_stream_replay_logs_once_when_polled_after_exhaustion(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: assert await _collect(iterator) == STREAM_EVENTS diff --git a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 8e439cc822f..281bb8f413b 100644 --- a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -14,7 +14,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, UnloadableEntitlementError, _agent_capped_servers, - _is_mcp_admitted_user_subject, + is_mcp_admitted_user_subject, ) from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -364,7 +364,7 @@ class TestMCPRequestHandler: "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", mock_manager, ), - patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), + patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), ): result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) @@ -385,7 +385,7 @@ class TestMCPRequestHandler: "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", mock_manager, ), - patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), + patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), ): result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) @@ -405,7 +405,7 @@ class TestMCPRequestHandler: "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", mock_manager, ), - patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), + patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[])), patch.object(MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=[])), ): @@ -430,7 +430,7 @@ class TestMCPRequestHandler: "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", mock_manager, ), - patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), + patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[])), patch.object( MCPRequestHandler, "_get_allowed_mcp_servers_for_team", @@ -733,7 +733,7 @@ class TestMCPRequestHandler: mock_manager, ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), ): servers = await MCPRequestHandler._team_granted_servers(team_obj, []) @@ -763,7 +763,7 @@ class TestMCPRequestHandler: "litellm.proxy.auth.auth_checks.get_team_object", AsyncMock(return_value=team_obj) ), patch( # test-quality-ok: access-group lookup hits the DB, not under test here - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[]), ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests @@ -801,7 +801,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=[]) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -838,7 +838,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=[]) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -880,7 +880,7 @@ class TestMCPRequestHandler: "litellm.proxy.auth.auth_checks.get_team_object", AsyncMock(return_value=team_obj) ), patch( # test-quality-ok: access-group lookup hits the DB, not under test here - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[]), ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests @@ -940,7 +940,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_org_object_permission", AsyncMock(return_value=org_object_permission) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -992,7 +992,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_user_object_permission", AsyncMock(return_value=user_object_permission) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -1045,7 +1045,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_org_object_permission", AsyncMock(return_value=org_object_permission) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -1067,7 +1067,7 @@ class TestMCPRequestHandler: MCPRequestHandler, "_get_user_object_permission", AsyncMock(return_value=user_object_permission) ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -1220,7 +1220,7 @@ class TestMCPRequestHandler: auth = UserAPIKeyAuth(api_key="k", access_group_ids=[]) with ( patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new=AsyncMock(return_value=[]), ), patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, @@ -1235,7 +1235,7 @@ class TestMCPRequestHandler: auth = UserAPIKeyAuth(api_key="k", access_group_ids=["grp-mcp"]) with ( patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new=AsyncMock(return_value=["alias-a", "srv-b"]), ), patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, @@ -1249,7 +1249,7 @@ class TestMCPRequestHandler: """Resolution failures degrade to no grants rather than raising.""" auth = UserAPIKeyAuth(api_key="k", access_group_ids=["grp-mcp"]) with patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new=AsyncMock(side_effect=Exception("db down")), ): result = await MCPRequestHandler._get_key_access_group_mcp_server_extras(auth) @@ -1720,7 +1720,7 @@ class TestMCPOAuth2AuthFlow: [b"sk-litellm-valid-key", b"Bearer sk-litellm-valid-key", b"bearer sk-litellm-valid-key"], ) async def test_x_litellm_api_key_survives_bearer_only_strip(self, header_value): - from litellm.proxy.auth.user_api_key_auth import _get_bearer_token + from litellm.proxy.auth.user_api_key_auth import get_bearer_token scope = { "type": "http", @@ -1741,7 +1741,7 @@ class TestMCPOAuth2AuthFlow: auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) mock_auth.assert_called_once() - assert _get_bearer_token(api_key=mock_auth.call_args.kwargs["api_key"]) == "sk-litellm-valid-key" + assert get_bearer_token(api_key=mock_auth.call_args.kwargs["api_key"]) == "sk-litellm-valid-key" assert auth_result.user_id == "test-user" async def test_litellm_key_in_authorization_backward_compat(self): @@ -4135,7 +4135,7 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper(): ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=["group-server1", "group-server2"], ) as mock_get_access_group_servers, @@ -4311,7 +4311,7 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission(): "litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock, ) as mock_get_perm: - with patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups") as mock_access_groups: + with patch.object(MCPRequestHandler, "get_mcp_servers_from_access_groups") as mock_access_groups: mock_access_groups.return_value = ["group-server"] result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth) @@ -4706,7 +4706,7 @@ class TestAgentMCPPermissions: mock_manager, ), patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here - MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) + MCPRequestHandler, "get_mcp_servers_from_access_groups", AsyncMock(return_value=[]) ), ) @@ -4944,7 +4944,7 @@ async def test_tool_permission_servers_included_in_allowed_servers(): patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=perm), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -5095,7 +5095,7 @@ class TestOrgMCPPermissions: ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -5121,7 +5121,7 @@ class TestOrgMCPPermissions: ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=["group_server_1"], ), @@ -5147,7 +5147,7 @@ class TestOrgMCPPermissions: ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -5552,7 +5552,7 @@ async def test_team_access_group_ids_resolve_to_mcp_servers(): return_value=mock_team, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-stripe"], ) as mock_resolver, @@ -5611,7 +5611,7 @@ async def test_team_access_group_ids_union_with_object_permission(): return_value=mock_team, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-stripe"], ), @@ -5649,7 +5649,7 @@ async def test_team_access_group_ids_empty_returns_no_extras(): return_value=mock_team, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=[], ) as mock_resolver, @@ -5712,7 +5712,7 @@ async def test_allowed_mcp_servers_for_key_excludes_access_group_ids(): with ( patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-stripe"], ) as mock_resolver, @@ -5760,7 +5760,7 @@ async def test_allowed_mcp_servers_for_key_uses_object_permission_not_access_gro with ( patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-stripe"], ) as mock_resolver, @@ -5787,7 +5787,7 @@ async def test_get_allowed_mcp_servers_surfaces_ungated_key_access_group_grant_e patches = _patch_proxy_server_globals_for_mcp() + [ patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=["srv-deepwiki"], ), @@ -5887,7 +5887,7 @@ async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynam return_value=team_obj, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -6022,13 +6022,13 @@ async def test_get_allowed_mcp_servers_team_all_proxy_key_scoped_to_one_end_to_e return_value=team_obj, ), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", new_callable=AsyncMock, return_value=[], ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -6737,7 +6737,7 @@ class TestMCPDcrBridgeDelegateAdmission: assert exc_info.value.status_code == 401 - _POLICY_GATE = "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp._run_centralized_common_checks" + _POLICY_GATE = "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.run_centralized_common_checks" async def _enforce_with_gate_error(self, error): """Drive _enforce_admitted_live_policy with the centralized gate raising ``error`` and return @@ -8273,7 +8273,7 @@ class TestUserSubjectTeamUnion: patch("litellm.proxy.auth.auth_checks.get_team_object", _get_team_object), patch("litellm.proxy.auth.auth_checks.get_user_object", _get_user_object), patch("litellm.proxy.auth.auth_checks.get_org_object", _get_org_object), - patch("litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[])), + patch("litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[])), patch("litellm.proxy.proxy_server.get_current_spend", _spend_from_fallback), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), @@ -8460,7 +8460,7 @@ class TestUserSubjectTeamUnion: one cross-team user drain several teams' buckets on a single call, blocking their other members for access those teams did not provide. Exactly one source is charged, and it is the SAME source billing picks — one owner for both, so they cannot disagree.""" - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 t1 = _make_team("t1", ["srv1"]) t1.metadata = {"mcp_rpm_limit": {"srv1": 5}} @@ -8474,7 +8474,7 @@ class TestUserSubjectTeamUnion: assert billed is not None and billed.team_id == "t1", "throttling and billing pick the same source" descriptors: list = [] - limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) + limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) limiter._add_mcp_per_team_rate_limit_descriptor(auth, "srv1", descriptors) charged = {d["value"]: d["rate_limit"]["requests_per_unit"] for d in descriptors} assert charged == {"t1:srv1": 5}, "only the attributing team's bucket is charged" @@ -8536,7 +8536,7 @@ class TestUserSubjectTeamUnion: server = MagicMock(server_id="srv1") with self._patch(teams_by_id={"t-grant": t_grant}, user_teams=["t-grant"]): with patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager._get_mcp_server_from_tool_name", + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_mcp_server_from_tool_name", MagicMock(return_value=server), ): billed = await MCPRequestHandler.billing_auth_for_tool_call(auth, tool_name="t-grant/tool_a") @@ -8985,7 +8985,7 @@ class TestUserSubjectTeamUnion: api_key="sk-real-key", metadata={"mcp_admitted_user_subject": True}, # caller-forged marker in key metadata ) - assert _is_mcp_admitted_user_subject(forged) is False + assert is_mcp_admitted_user_subject(forged) is False with self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]): assert await MCPRequestHandler._team_ids_for_mcp_grant(forged) == [] assert await MCPRequestHandler._get_allowed_mcp_servers_for_team(forged) == [] @@ -9065,8 +9065,8 @@ class TestUserSubjectTeamUnion: via_validate = UserAPIKeyAuth.model_validate({"user_id": "u", "mcp_admitted_user_subject": True}) assert via_kwarg.mcp_admitted_user_subject is False assert via_validate.mcp_admitted_user_subject is False - assert _is_mcp_admitted_user_subject(via_kwarg) is False - assert _is_mcp_admitted_user_subject(via_validate) is False + assert is_mcp_admitted_user_subject(via_kwarg) is False + assert is_mcp_admitted_user_subject(via_validate) is False @pytest.mark.asyncio @@ -9126,7 +9126,7 @@ class TestAdmittedSubjectPerTeamOrgCap: patch("litellm.proxy.auth.auth_checks.get_org_object", _get_org_object), patch("litellm.proxy.auth.auth_checks.get_object_permission", _get_object_permission), patch( - "litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_mcp_server_ids_from_access_groups", AsyncMock(return_value=[]), ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), @@ -9574,7 +9574,7 @@ class TestUserMCPEntitlement: ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -9731,7 +9731,7 @@ class TestUserMCPEntitlement: with self._entitled(self._perm(tool_permissions={"srv-a": ["read"]})): with patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py index f57b5121fa5..de86c1b62f4 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py @@ -1541,7 +1541,7 @@ async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch): ) ) - monkeypatch.setattr(enc, "_get_salt_key", lambda: key_old) + monkeypatch.setattr(enc, "get_salt_key", lambda: key_old) prisma = MagicMock() prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) @@ -1558,7 +1558,7 @@ async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch): assert store_update.await_args.kwargs["where"] == {"server_id": "config_faros"} rotated_blob = store_update.await_args.kwargs["data"]["credentials"] - monkeypatch.setattr(enc, "_get_salt_key", lambda: key_new) + monkeypatch.setattr(enc, "get_salt_key", lambda: key_new) recovered = decrypt_credentials(credentials=json.loads(rotated_blob)) assert recovered["client_id"] == "cid-123" assert recovered["client_secret"] == "sec-456" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index bf21e3434ca..23bdfa3eadf 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -251,7 +251,7 @@ async def test_register_resolves_cold_oauth_metadata(): "_discover_oauth_metadata_for_server", new=AsyncMock(return_value=_resolved_oauth_metadata()), ) as discovery, - patch.object(discoverable_endpoints, "_read_request_body", new=AsyncMock(return_value={})), + patch.object(discoverable_endpoints, "read_request_body", new=AsyncMock(return_value={})), patch.object( discoverable_endpoints, "get_async_httpx_client", @@ -307,7 +307,7 @@ async def test_register_route_bridge_missing_registration_url_joins_discovery(): ) as discovery, patch.object( # test-quality-ok: the MagicMock Request carries no body; this seam feeds the RFC 7591 redirect_uris discoverable_endpoints, - "_read_request_body", + "read_request_body", new=AsyncMock(return_value={"redirect_uris": ["https://client.example.com/cb"]}), ), patch.object( # test-quality-ok: keeps the DCR POST off the network so its target URL can be asserted @@ -766,7 +766,7 @@ async def test_register_client_without_mcp_server_name_returns_dummy(server_name mock_request.base_url = "https://proxy.litellm.example/" mock_request.headers = {} with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={}), ): result = await register_client(request=mock_request, mcp_server_name=server_name) @@ -817,7 +817,7 @@ async def test_register_client_returns_existing_server_credentials(use_root): try: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={}), ): result = await register_client( @@ -889,7 +889,7 @@ async def test_register_client_remote_registration_success(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value=request_payload), ), patch( @@ -975,7 +975,7 @@ async def test_register_client_non_bridge_returns_client_redirect_not_gateway_ca try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value=request_payload), ), patch( @@ -1031,7 +1031,7 @@ async def test_register_client_admin_client_id_echoes_client_redirect_uris(): try: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": [client_redirect]}), ): result = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name) @@ -1107,7 +1107,7 @@ async def test_dcr_full_loop_lands_on_client_redirect_not_gateway_callback(monke try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock( return_value={ "client_name": "Open WebUI", @@ -1268,7 +1268,7 @@ async def test_register_client_malformed_redirect_uris_falls_back_to_gateway_cal try: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": malformed_redirect_uris}), ): result = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name) @@ -1312,7 +1312,7 @@ async def test_register_client_valid_multi_redirect_uris_all_echoed(): client_redirects = ["https://app.example/cb", "http://127.0.0.1:6274/callback"] try: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": client_redirects}), ): result = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name) @@ -2258,7 +2258,7 @@ async def test_register_client_reuses_existing_client_id_without_re_dcr(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value=request_payload), ), patch( @@ -2337,7 +2337,7 @@ async def test_public_register_route_does_not_persist_client_credentials(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value=request_payload), ), patch( @@ -2631,7 +2631,7 @@ async def test_register_client_respects_x_forwarded_proto(): mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={}), ): result = await register_client(request=mock_request) @@ -3737,7 +3737,7 @@ async def test_register_client_resolves_server_by_id_when_name_lookup_fails(): with ( patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam - patch.object(discoverable_endpoints, "_read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: request seam + patch.object(discoverable_endpoints, "read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: request seam patch.object(discoverable_endpoints, "get_async_httpx_client", return_value=client), # test-quality-ok: HTTP seam ): result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id) @@ -4065,7 +4065,7 @@ async def test_register_root_does_aggregate_dcr_not_single_server_resolution(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}), ), patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"), @@ -4106,7 +4106,7 @@ async def test_register_root_does_not_leak_a_private_server(): try: with ( patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}), ), patch( @@ -5879,7 +5879,7 @@ async def test_interactive_bridge_authorize_seals_sso_user_into_state(): with ( patch( - "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.user_id_from_session_cookie", return_value="sso-user-42", ), patch( @@ -5944,7 +5944,7 @@ async def test_bridge_authorize_gates_on_the_egress_server_access_resolver(user_ try: with ( patch( - "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.user_id_from_session_cookie", return_value="bridge-user-1", ), patch( @@ -6015,7 +6015,7 @@ async def test_bridge_authorize_reload_failure_denies_or_stays_retryable(reload_ try: with ( patch( - "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.user_id_from_session_cookie", return_value="bridge-user-1", ), patch( @@ -6061,7 +6061,7 @@ async def test_interactive_bridge_authorize_without_session_redirects_to_login() server = _bridge_server(auth_type=MCPAuth.oauth_delegate, client_id="admin-client", registration_url=None) with patch( - "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints.user_id_from_session_cookie", return_value=None, ): response = await authorize_with_server( @@ -6600,7 +6600,7 @@ async def test_revalidate_active_subject_dispatches_on_subject_type(): new=AsyncMock(return_value=_ResolvedKey(key_hash="kh", key=MagicMock())), ) as key_reload, patch( - "litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_user_by_id", + "litellm.proxy._experimental.mcp_server.bridge_token_flow.reload_active_user_by_id", new=AsyncMock(return_value=None), ) as user_reload, ): @@ -6614,7 +6614,7 @@ async def test_revalidate_active_subject_dispatches_on_subject_type(): new=AsyncMock(), ) as key_reload2, patch( - "litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_user_by_id", + "litellm.proxy._experimental.mcp_server.bridge_token_flow.reload_active_user_by_id", new=AsyncMock(return_value="no_active_key"), ) as user_reload2, ): @@ -7262,7 +7262,7 @@ def test_bridge_reported_expires_in_can_be_zero_at_jwt_exp_boundary(): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( _BridgeMintReady, - _finish_bridge_mint, + finish_bridge_mint, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( envelope_keys_from_master_key, @@ -7274,7 +7274,7 @@ def test_bridge_reported_expires_in_can_be_zero_at_jwt_exp_boundary(): identity=key_hash_identity(server_id="bridge_srv", key_hash="hashed-litellm-key-77"), keys=envelope_keys_from_master_key(_BRIDGE_MASTER_KEY), ) - response = _finish_bridge_mint( + response = finish_bridge_mint( ready=ready, mcp_server=_bridge_server(auth_type=MCPAuth.oauth_delegate), token_response={"access_token": "UP", "expires_in": 1}, @@ -7439,7 +7439,7 @@ async def _exchange_persistence_attempted_for_auth_type(auth_type) -> bool: return_value=fake_http_client, ), patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.extract_user_id_from_request", new_callable=AsyncMock, return_value="admin-user", ), @@ -7623,7 +7623,7 @@ async def test_extract_user_id_reads_x_litellm_api_key_header(proxy_globals): Authorization. Reading only Authorization dropped the identity, so the per-user token was never stored and the egress 401'd forever. Resolution must honor x-litellm-api-key.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import UserAPIKeyAuth, hash_token from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7639,7 +7639,7 @@ async def test_extract_user_id_reads_x_litellm_api_key_header(proxy_globals): proxy_globals.prisma_client = object() request = _token_request({"x-litellm-api-key": f"Bearer {key}"}) - assert await _extract_user_id_from_request(request) == "alice" + assert await extract_user_id_from_request(request) == "alice" @pytest.mark.asyncio @@ -7648,7 +7648,7 @@ async def test_extract_user_id_rehydrates_cross_replica_dict_cache(proxy_globals Resolution must rehydrate it; the old getattr(cached, "user_id") returned None on a dict, which is exactly why a multi-replica gateway never found the stored token.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import hash_token from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7660,7 +7660,7 @@ async def test_extract_user_id_rehydrates_cross_replica_dict_cache(proxy_globals proxy_globals.prisma_client = object() request = _token_request({"Authorization": f"Bearer {key}"}) - assert await _extract_user_id_from_request(request) == "alice" + assert await extract_user_id_from_request(request) == "alice" @pytest.mark.asyncio @@ -7669,7 +7669,7 @@ async def test_extract_user_id_falls_back_to_db_on_cache_miss(proxy_globals): cache-only peek and skipped the DB, so any replica that hadn't just authenticated the key failed to store the token.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7684,14 +7684,14 @@ async def test_extract_user_id_falls_back_to_db_on_cache_miss(proxy_globals): proxy_globals.prisma_client = _FakePrisma() request = _token_request({"x-litellm-api-key": key}) - assert await _extract_user_id_from_request(request) == "db-bob" + assert await extract_user_id_from_request(request) == "db-bob" @pytest.mark.asyncio async def test_extract_user_id_none_without_litellm_key(proxy_globals): """No LiteLLM key on the request resolves to None without consulting the resolver.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7699,7 +7699,7 @@ async def test_extract_user_id_none_without_litellm_key(proxy_globals): proxy_globals.prisma_client = object() request = _token_request({"content-type": "application/json"}) - assert await _extract_user_id_from_request(request) is None + assert await extract_user_id_from_request(request) is None @pytest.mark.asyncio @@ -7708,7 +7708,7 @@ async def test_extract_user_id_rejects_blocked_key(proxy_globals): checking blocked/expiry (the main auth pipeline does, and the public token endpoint bypasses it), so a revoked key could otherwise overwrite the stored per-user OAuth token for its user.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7721,7 +7721,7 @@ async def test_extract_user_id_rejects_blocked_key(proxy_globals): proxy_globals.prisma_client = _FakePrisma() request = _token_request({"x-litellm-api-key": "sk-blocked-key"}) - assert await _extract_user_id_from_request(request) is None + assert await extract_user_id_from_request(request) is None @pytest.mark.asyncio @@ -7730,7 +7730,7 @@ async def test_extract_user_id_rejects_expired_key(proxy_globals): from datetime import datetime, timedelta, timezone from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -7745,7 +7745,7 @@ async def test_extract_user_id_rejects_expired_key(proxy_globals): proxy_globals.prisma_client = _FakePrisma() request = _token_request({"x-litellm-api-key": "sk-expired-key"}) - assert await _extract_user_id_from_request(request) is None + assert await extract_user_id_from_request(request) is None @pytest.mark.asyncio @@ -7785,7 +7785,7 @@ async def test_resolve_active_litellm_key_resolves_key_without_user_id(proxy_glo blocked and expiry, and the key hash (not the user) is what the mint seals. The per-user token store still gets no user for such a key, since there is none to key a stored credential by.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, + extract_user_id_from_request, _ResolvedKey, _resolve_active_litellm_key, ) @@ -7806,7 +7806,7 @@ async def test_resolve_active_litellm_key_resolves_key_without_user_id(proxy_glo resolved = await _resolve_active_litellm_key(request) assert isinstance(resolved, _ResolvedKey) assert resolved.key_hash == hash_token(key) - assert await _extract_user_id_from_request(request) is None + assert await extract_user_id_from_request(request) is None @pytest.mark.asyncio @@ -7979,7 +7979,7 @@ async def test_reload_active_user_by_id_missing_user_is_no_active_key(proxy_glob refresh path maps it to invalid_grant), not unresolvable/500. get_user_object catches the missing row and re-raises a bare ValueError, so a missing user must not be misclassified as a DB outage or an opaque gateway fault.""" - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id + from litellm.proxy._experimental.mcp_server.bridge_token_flow import reload_active_user_by_id from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache proxy_globals.user_api_key_cache = UserApiKeyCache() @@ -7989,7 +7989,7 @@ async def test_reload_active_user_by_id_missing_user_is_no_active_key(proxy_glob "litellm.proxy.auth.auth_checks.get_user_object", new=AsyncMock(side_effect=_wrapped_user_lookup_error(Exception())), ): - assert await _reload_active_user_by_id("gone-user") == "no_active_key" + assert await reload_active_user_by_id("gone-user") == "no_active_key" @pytest.mark.asyncio @@ -7998,7 +7998,7 @@ async def test_reload_active_user_by_id_db_outage_is_unavailable(proxy_globals): a missing user, so the refresh path surfaces "unavailable" (a 503) rather than blaming the caller. get_user_object wraps the outage in a bare ValueError, so this exercises the chain-aware classifier; a raw ConnectionError would falsely pass even a chain-blind check because it is an OSError.""" - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id + from litellm.proxy._experimental.mcp_server.bridge_token_flow import reload_active_user_by_id from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache proxy_globals.user_api_key_cache = UserApiKeyCache() @@ -8008,7 +8008,7 @@ async def test_reload_active_user_by_id_db_outage_is_unavailable(proxy_globals): "litellm.proxy.auth.auth_checks.get_user_object", new=AsyncMock(side_effect=_wrapped_user_lookup_error(ConnectionError("user database unreachable"))), ): - assert await _reload_active_user_by_id("sso-user-7") == "unavailable" + assert await reload_active_user_by_id("sso-user-7") == "unavailable" @pytest.mark.asyncio @@ -8018,7 +8018,7 @@ async def test_reload_active_user_by_id_permanent_engine_fault_is_faulted(proxy_ wraps the fault in a bare ValueError, so the classification has to read the wrapped cause.""" from prisma.engine.errors import MismatchedVersionsError - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id + from litellm.proxy._experimental.mcp_server.bridge_token_flow import reload_active_user_by_id from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache proxy_globals.user_api_key_cache = UserApiKeyCache() @@ -8028,7 +8028,7 @@ async def test_reload_active_user_by_id_permanent_engine_fault_is_faulted(proxy_ "litellm.proxy.auth.auth_checks.get_user_object", new=AsyncMock(side_effect=_wrapped_user_lookup_error(MismatchedVersionsError(expected="1", got="2"))), ): - assert await _reload_active_user_by_id("sso-user-7") == "faulted" + assert await reload_active_user_by_id("sso-user-7") == "faulted" @pytest.mark.asyncio @@ -8067,7 +8067,7 @@ async def test_load_active_user_by_id_serves_a_cached_row_without_a_database_rea cached row answers without a database read, and only a caller that asks for the database row pays for one.""" from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _reload_active_user_by_id, + reload_active_user_by_id, load_active_user_by_id, ) from litellm.proxy._types import LiteLLM_UserTable @@ -8090,7 +8090,7 @@ async def test_load_active_user_by_id_serves_a_cached_row_without_a_database_rea assert not isinstance(loaded, str) assert loaded.teams == ["team-a"] - assert await _reload_active_user_by_id("cached-jwt-user") is None + assert await reload_active_user_by_id("cached-jwt-user") is None prisma.db.litellm_usertable.find_unique.assert_not_awaited() @@ -8343,7 +8343,7 @@ async def test_register_client_rejects_non_oauth2_server(): try: with pytest.raises(HTTPException) as exc_info: with patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.read_request_body", new=AsyncMock(return_value={}), ): await register_client(request=mock_request, mcp_server_name="access_group_server") @@ -8731,7 +8731,7 @@ async def test_token_exchange_refresh_passes_presented_refresh_ownership(): return_value=client, ), patch( # test-quality-ok: no injection seam exists for request identity extraction - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.extract_user_id_from_request", new=AsyncMock(return_value="user-a"), ), patch( # test-quality-ok: captures the ownership value at the exchange boundary @@ -8797,7 +8797,7 @@ async def test_token_exchange_authorization_code_passes_no_refresh_ownership(mon return_value=client, ), patch( # test-quality-ok: no injection seam exists for request identity extraction - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.extract_user_id_from_request", new=AsyncMock(return_value="user-a"), ), patch( # test-quality-ok: captures the ownership value at the exchange boundary @@ -9294,7 +9294,7 @@ async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch, lega issuer="https://idp.example", ) - monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-hydrate-key") + monkeypatch.setattr(enc, "get_salt_key", lambda: "salt-hydrate-key") stored_blob = safe_dumps( encrypt_credentials( credentials={**({} if legacy else {"dcr_issuer": "https://idp.example", "dcr_server_url": "https://resource.example/mcp"}), @@ -9351,7 +9351,7 @@ async def test_reuse_config_server_reads_store_with_real_crypto(monkeypatch): issuer="https://idp.example", ) - monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-reuse-key") + monkeypatch.setattr(enc, "get_salt_key", lambda: "salt-reuse-key") blob = safe_dumps( encrypt_credentials( credentials={"dcr_issuer": "https://idp.example", "dcr_server_url": "https://resource.example/mcp","client_id": "stored-client", "client_secret": "sec", "redirect_uris": ["https://x/callback"]}, @@ -11697,7 +11697,7 @@ def test_introspect_route_answers_for_authenticated_caller(monkeypatch): return None monkeypatch.setattr( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._reload_active_user_by_id", fake_reload + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.reload_active_user_by_id", fake_reload ) app = FastAPI() app.include_router(router) @@ -12153,7 +12153,7 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( from litellm.proxy import proxy_server from litellm.proxy._experimental.mcp_server import mcp_server_manager from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( - _extract_user_id_from_request, authorize_oauth_credential_request, + extract_user_id_from_request, authorize_oauth_credential_request, ) allowed_servers: Final = AsyncMock(return_value=["server-a"]) @@ -12184,7 +12184,7 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( request: Final = _token_request({"Authorization": f"Bearer {bearer}"}) result: Final = ( await authorize_oauth_credential_request(request, "server-a") - if credential_write else await _extract_user_id_from_request(request) + if credential_write else await extract_user_id_from_request(request) ) assert result is None allowed_servers.assert_not_awaited() @@ -12196,7 +12196,7 @@ async def test_oauth_jwt_cannot_override_explicit_litellm_key( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], blocked: bool, ) -> None: - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request from litellm.proxy._types import UserAPIKeyAuth, hash_token handler, signing_key = jwt_oauth_identity @@ -12208,7 +12208,7 @@ async def test_oauth_jwt_cannot_override_explicit_litellm_key( "x-litellm-api-key": key, } ) - assert await _extract_user_id_from_request(request) == (None if blocked else "key-owner") + assert await extract_user_id_from_request(request) == (None if blocked else "key-owner") @pytest.mark.asyncio @@ -12220,7 +12220,7 @@ async def test_oauth_jwt_uses_configured_virtual_key_owner( mapping: str, ) -> None: from litellm.models.user import LiteLLM_UserTable - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request from litellm.proxy._types import UserAPIKeyAuth, UnregisteredJWTClientBehavior, hash_token from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key @@ -12248,7 +12248,7 @@ async def test_oauth_jwt_uses_configured_virtual_key_owner( ) request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}) expected: Final = "jwt-owner" if mapping == "fallback" else "mapped-owner" if mapping == "active" else None - assert await _extract_user_id_from_request(request) == expected + assert await extract_user_id_from_request(request) == expected @pytest.mark.asyncio @@ -12257,13 +12257,13 @@ async def test_oauth_jwt_respects_custom_validation_and_email_policy( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], allowed_domain: str | None, ) -> None: - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request handler, signing_key = jwt_oauth_identity handler.litellm_jwtauth.custom_validate = lambda claims: True handler.litellm_jwtauth.user_allowed_email_domain = allowed_domain request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}) - assert await _extract_user_id_from_request(request) == (None if allowed_domain else "jwt-owner") + assert await extract_user_id_from_request(request) == (None if allowed_domain else "jwt-owner") @pytest.mark.asyncio @@ -12274,7 +12274,7 @@ async def test_oauth_jwt_identity_preserves_separate_mcp_route_authorization( route_allowed: bool, ) -> None: from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request from litellm.proxy._types import LitellmUserRoles, RoleBasedPermissions, RoleMapping from litellm.proxy.auth.handle_jwt import JWTAuthManager @@ -12301,7 +12301,7 @@ async def test_oauth_jwt_identity_preserves_separate_mcp_route_authorization( ) bearer: Final = _oauth_identity_jwt(signing_key) request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path="/example/token") - assert await _extract_user_id_from_request(request) == "jwt-owner" + assert await extract_user_id_from_request(request) == "jwt-owner" admission: Final = JWTAuthManager.auth_builder( api_key=bearer, jwt_handler=handler, @@ -12335,7 +12335,7 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( ) -> None: from litellm.models.user import LiteLLM_UserTable from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request from litellm.proxy.auth.handle_jwt import JWTAuthManager handler, signing_key = jwt_oauth_identity @@ -12356,7 +12356,7 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( monkeypatch.setattr(proxy_server, "prisma_client", database) bearer: Final = _oauth_identity_jwt(signing_key, owner=external_id, scope="litellm_proxy_admin" if admin else "") request: Final = _token_request({"Authorization": f"Bearer {bearer}"}) - stored_owner: Final = await _extract_user_id_from_request(request) + stored_owner: Final = await extract_user_id_from_request(request) assert stored_owner == (None if inactive else external_id if admin else "canonical-oauth-owner") assert table.find_unique.await_count == 2 if identity == "email": @@ -12382,7 +12382,7 @@ async def test_oauth_jwt_identity_does_not_provision_or_synchronize_teams( ) -> None: from litellm.models.user import LiteLLM_UserTable from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import extract_user_id_from_request handler, signing_key = jwt_oauth_identity handler.litellm_jwtauth.enforce_team_based_model_access = True @@ -12394,7 +12394,7 @@ async def test_oauth_jwt_identity_does_not_provision_or_synchronize_teams( request: Final = _token_request( {"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}, path="/example/token" ) - assert await _extract_user_id_from_request(request) == "jwt-owner" + assert await extract_user_id_from_request(request) == "jwt-owner" assert owner.teams == ["existing-team"] proxy_server.prisma_client.db.litellm_teamtable.find_unique.assert_not_called() proxy_server.prisma_client.db.litellm_teamtable.upsert.assert_not_called() @@ -12410,7 +12410,7 @@ async def test_oauth_refresh_revalidates_the_same_active_user_rule( ) -> None: from litellm.models.user import LiteLLM_UserTable from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id + from litellm.proxy._experimental.mcp_server.bridge_token_flow import reload_active_user_by_id handler, _ = jwt_oauth_identity user_id: Final = f"jwt-owner-{state}" @@ -12419,7 +12419,7 @@ async def test_oauth_refresh_revalidates_the_same_active_user_rule( if state == "missing_database": monkeypatch.setattr(proxy_server, "prisma_client", None) expected: Final = None if state == "active" else "no_active_key" if state == "inactive" else "unresolvable" - assert await _reload_active_user_by_id(user_id) == expected + assert await reload_active_user_by_id(user_id) == expected if state != "missing_database": cached: Final = handler.user_api_key_cache.get_cache(user_id, model_type=LiteLLM_UserTable) assert cached is not None diff --git a/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py b/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py index d8e4a342e52..eddfdc4f884 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py @@ -308,7 +308,7 @@ async def test_e2e_jwt_team_mcp_permissions_enforced(monkeypatch): # Mock _get_mcp_servers_from_access_groups to return empty with patch.object( - MCPRequestHandler, "_get_mcp_servers_from_access_groups" + MCPRequestHandler, "get_mcp_servers_from_access_groups" ) as mock_access_groups: mock_access_groups.return_value = [] @@ -506,7 +506,7 @@ async def test_e2e_jwt_team_mcp_key_intersection(monkeypatch): ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), diff --git a/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py b/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py index 052231b562a..7ee340a69f4 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py @@ -64,7 +64,7 @@ async def test_simple_jwt_mcp_permissions_enforced(): ), patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), @@ -143,7 +143,7 @@ async def test_simple_jwt_team_id_required_for_mcp_permissions(): ) as mock_get_team, patch.object( MCPRequestHandler, - "_get_mcp_servers_from_access_groups", + "get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ), diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py index 6a08e8eeab6..6970ffb6b08 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -631,7 +631,7 @@ async def test_health_check_reaches_servers_without_forwarding_per_user_env_vars manager: Final = MCPServerManager() manager.registry[mock_server.server_id] = mock_server create_client: Final = AsyncMock() - monkeypatch.setattr(manager, "_create_mcp_client", create_client) + monkeypatch.setattr(manager, "create_mcp_client", create_client) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") route: Final = respx_mock.get(mock_server.url).respond(401) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py index 1578ea8e601..8d5e36e9795 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py @@ -68,9 +68,9 @@ def _fake_proxy_logging(capture: dict, *, guardrail_effect=None): """ plo = mock.MagicMock() plo.enforce_mcp_server_rate_limits = mock.AsyncMock() - plo._create_mcp_request_object_from_kwargs.return_value = mock.MagicMock() + plo.create_mcp_request_object_from_kwargs.return_value = mock.MagicMock() # Mirror the real conversion's metadata bucket so a test can prove it survives. - plo._convert_mcp_to_llm_format.side_effect = lambda *_a, **_k: { + plo.convert_mcp_to_llm_format.side_effect = lambda *_a, **_k: { "metadata": {"headers": {"x-forwarded-for": "1.2.3.4"}} } diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 333984bde68..c477a8d34cb 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -99,10 +99,10 @@ class TestPreCallToolCheckReturnsHeaders: server = self._make_server() proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value={"modified_arguments": {"key": "val"}}) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {"key": "val"}}) + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {"key": "val"}}) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): with patch.object( @@ -130,10 +130,10 @@ class TestPreCallToolCheckReturnsHeaders: hook_headers = {"Authorization": "Bearer signed-jwt", "X-Trace-Id": "abc123"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value={"extra_headers": hook_headers}) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock( return_value={"arguments": {"key": "val"}, "extra_headers": hook_headers} ) @@ -161,8 +161,8 @@ class TestPreCallToolCheckReturnsHeaders: server = self._make_server() proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): @@ -193,10 +193,10 @@ class TestPreCallToolCheckReturnsHeaders: modified_args = {"key": "modified", "extra": "added"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value={"modified_arguments": modified_args}) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": modified_args}) + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": modified_args}) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): with patch.object( @@ -226,10 +226,10 @@ class TestPreCallToolCheckReturnsHeaders: hook_headers = {"Authorization": "Bearer jwt"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={"model": "fake"}) proxy_logging.pre_call_hook = AsyncMock(return_value={"dummy": True}) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock( return_value={"arguments": modified_args, "extra_headers": hook_headers} ) @@ -276,7 +276,7 @@ class TestCallToolFlowsHookHeaders: with patch.object( manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=server, ): with patch.object( @@ -317,7 +317,7 @@ class TestCallToolFlowsHookHeaders: with patch.object( manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=server, ): with patch.object( @@ -347,7 +347,7 @@ class TestCallToolFlowsHookHeaders: with patch.object( manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=server, ): with patch.object( @@ -394,7 +394,7 @@ class TestCallToolFlowsHookHeaders: spec_path="/path/to/spec.yaml", ) - with patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server): + with patch.object(manager, "get_mcp_server_from_tool_name", return_value=server): with patch.object( manager, "pre_call_tool_check", @@ -441,7 +441,7 @@ class TestCallToolFlowsHookHeaders: spec_path="/path/to/spec.yaml", ) - with patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server): + with patch.object(manager, "get_mcp_server_from_tool_name", return_value=server): with patch.object( manager, "pre_call_tool_check", @@ -504,8 +504,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -540,8 +540,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -584,8 +584,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -631,8 +631,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -667,8 +667,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -708,8 +708,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -752,8 +752,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -790,8 +790,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -828,8 +828,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -875,8 +875,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -919,8 +919,8 @@ class TestHookHeaderMergePriority: mock_client.call_tool = AsyncMock(return_value=MagicMock()) return mock_client - with patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client): - with patch.object(manager, "_build_stdio_env", return_value=None): + with patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client): + with patch.object(manager, "build_stdio_env", return_value=None): try: await manager._call_regular_mcp_tool( mcp_server=server, @@ -1031,10 +1031,10 @@ class TestMcpRateLimitServerNameSurfacing: return {"model": "fake"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(side_effect=capture_convert) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(side_effect=capture_convert) proxy_logging.pre_call_hook = AsyncMock(return_value=None) - proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {}}) + proxy_logging.convert_mcp_hook_response_to_kwargs = MagicMock(return_value={"arguments": {}}) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): with patch.object( @@ -1060,7 +1060,7 @@ class TestOpenApiByokCallTool: async def test_call_tool_openapi_byok_injects_request_auth_contextvar(self): """Playground/responses call call_tool directly; BYOK must reach OpenAPI handlers.""" from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, + request_auth_header, ) manager = MCPServerManager() @@ -1078,7 +1078,7 @@ class TestOpenApiByokCallTool: captured_auth: dict[str, Optional[str]] = {} async def fake_openapi_handler(_server, _name, _arguments, _wire_compat): - captured_auth["value"] = _request_auth_header.get() + captured_auth["value"] = request_auth_header.get() return CallToolResult(content=[], isError=False) with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server): @@ -1305,7 +1305,7 @@ class TestOpenApiResolvedUpstreamAuth: """The managed spec_path arm resolves the v2 credential and sets the ContextVar; kills the mutant that drops the resolve_openapi_upstream_auth call in call_tool.""" from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( StaticHeaderAuth, @@ -1318,7 +1318,7 @@ class TestOpenApiResolvedUpstreamAuth: captured: Dict[str, Any] = {} async def fake_openapi_handler(_server, _name, _arguments, _wire_compat): - captured["resolved"] = _request_resolved_auth_headers.get() + captured["resolved"] = request_resolved_auth_headers.get() return CallToolResult(content=[], isError=False) with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server): @@ -1336,7 +1336,7 @@ class TestOpenApiResolvedUpstreamAuth: ) assert captured["resolved"] == {"Authorization": "Bearer stored-user-token"} - assert _request_resolved_auth_headers.get() is None + assert request_resolved_auth_headers.get() is None @pytest.mark.asyncio async def test_call_tool_openapi_m2m_missing_token_url_fails_closed(self): @@ -1467,8 +1467,8 @@ class TestPreCallToolCheckExposesClientHeaders: return {"model": "fake"} proxy_logging = MagicMock(spec=ProxyLogging) - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) - proxy_logging._convert_mcp_to_llm_format = MagicMock(side_effect=capture) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock()) + proxy_logging.convert_mcp_to_llm_format = MagicMock(side_effect=capture) proxy_logging.pre_call_hook = AsyncMock(return_value=None) with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py index 0dc7ac5ecd9..d046871c8e2 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py @@ -57,7 +57,7 @@ def _patch_client_with_tracker(manager: MCPServerManager, tracker: _ConcurrencyT return _ProbeClient() - return patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client) + return patch.object(manager, "create_mcp_client", side_effect=fake_create_mcp_client) async def _fire(manager: MCPServerManager, server: MCPServer, n: int) -> None: diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py index 1909e3306a2..7290f0de86d 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -347,7 +347,7 @@ async def test_aggregate_list_tools_absorbs_one_unauthenticated_server(): ), patch.object( mcp_operations, "filter_tools_by_key_team_permissions", AsyncMock(side_effect=lambda tools, **k: tools) ), patch.object( - mcp_operations.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) + mcp_operations.global_mcp_server_manager, "get_tools_from_server", AsyncMock(side_effect=fake_get_tools) ): listing = await mcp_operations._get_tools_from_mcp_servers( user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"), @@ -370,7 +370,7 @@ async def test_single_server_route_also_absorbs_upstream_auth_error(): from unittest.mock import patch from litellm.proxy._experimental.mcp_server import server as mcp_server - from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_gateway_server_name + from litellm.proxy._experimental.mcp_server.mcp_context import mcp_gateway_server_name from litellm.proxy._types import UserAPIKeyAuth delegate = _http_server( @@ -381,14 +381,14 @@ async def test_single_server_route_also_absorbs_upstream_auth_error(): raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name) # //mcp sets the path-derived single-server scope; absorption must hold even then. - token = _mcp_gateway_server_name.set("delegate_docs") + token = mcp_gateway_server_name.set("delegate_docs") try: with patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object( mcp_operations, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={}) ), patch.object(mcp_operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( mcp_operations, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) ), patch.object( - mcp_operations.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) + mcp_operations.global_mcp_server_manager, "get_tools_from_server", AsyncMock(side_effect=fake_get_tools) ): listing = await mcp_operations._get_tools_from_mcp_servers( user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"), @@ -398,7 +398,7 @@ async def test_single_server_route_also_absorbs_upstream_auth_error(): assert listing.tools == [] assert listing.outcomes["delegate_docs"].tag == "auth_required" finally: - _mcp_gateway_server_name.reset(token) + mcp_gateway_server_name.reset(token) @pytest.mark.asyncio @@ -425,7 +425,7 @@ async def test_aggregate_with_single_accessible_server_still_absorbs(): ), patch.object(mcp_operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( mcp_operations, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) ), patch.object( - mcp_operations.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) + mcp_operations.global_mcp_server_manager, "get_tools_from_server", AsyncMock(side_effect=fake_get_tools) ): # Aggregate route: no explicit server filter, even though only one server is accessible. listing = await mcp_operations._get_tools_from_mcp_servers( @@ -470,7 +470,7 @@ async def test_client_creation_failure_logs_sanitized_exchange(monkeypatch, capl request = httpx.Request("POST", "https://upstream/mcp?credential=query-secret") response = httpx.Response(500, request=request, json={"error":"missing_scope"}) error = httpx.HTTPStatusError("query-secret", request=request, response=response) - monkeypatch.setattr(manager, "_create_mcp_client", AsyncMock(side_effect=error)) + monkeypatch.setattr(manager, "create_mcp_client", AsyncMock(side_effect=error)) with caplog.at_level(logging.WARNING, logger="LiteLLM"): with pytest.raises(MCPServerListError): await manager._get_tools_from_server(server) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py index e52a86d76af..6f1bbad5040 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py @@ -164,7 +164,7 @@ async def test_proxy_call_tool_on_a_never_listed_tool_hands_the_pre_hook_no_list patch.dict(manager.registry, {server.server_id: server}), patch.dict(manager.tool_name_to_mcp_server_name_mapping), patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), - patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "create_mcp_client", AsyncMock(return_value=object())), patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), patch.object(manager, "pre_call_tool_check", pre_call_tool_check), patch.object(manager, "_call_regular_mcp_tool", call_regular_mcp_tool), diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index 941d67cb587..23defa6f513 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -925,7 +925,7 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_mcp_server_by_id = lambda server_id: ( mock_server_1 if server_id == "server1_id" else mock_server_2 ) - mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( side_effect=lambda server_ids, client_ip: (server_ids, 0) @@ -972,7 +972,7 @@ async def test_get_tools_from_mcp_servers(): return [mock_tool_1] return [mock_tool_2] - mock_manager_2._get_tools_from_server = AsyncMock( + mock_manager_2.get_tools_from_server = AsyncMock( side_effect=mock_get_tools_side_effect ) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) @@ -1007,7 +1007,7 @@ async def test_get_tools_from_mcp_servers(): if server_id == "server1_id" else (mock_server_2 if server_id == "server2_id" else mock_server_3) ) - mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( side_effect=lambda server_ids, client_ip: (server_ids, 0) @@ -1018,7 +1018,7 @@ async def test_get_tools_from_mcp_servers(): mock_manager, ): with patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", AsyncMock(return_value=["server3_id"]), ): # Test with specific servers @@ -1399,7 +1399,7 @@ async def test_mcp_server_manager_alias_tool_prefixing(): mock_client_constructor, ): # Get tools from server - tools = await test_manager._get_tools_from_server(mock_server) + tools = await test_manager.get_tools_from_server(mock_server) # Verify tool is prefixed with alias assert len(tools) == 1 @@ -1459,7 +1459,7 @@ async def test_mcp_server_manager_server_name_tool_prefixing(): mock_client_constructor, ): # Get tools from server - tools = await test_manager._get_tools_from_server(mock_server) + tools = await test_manager.get_tools_from_server(mock_server) # Verify tool is prefixed with server_name (normalized) assert len(tools) == 1 @@ -1519,7 +1519,7 @@ async def test_mcp_server_manager_server_id_tool_prefixing(): mock_client_constructor, ): # Get tools from server - tools = await test_manager._get_tools_from_server(mock_server) + tools = await test_manager.get_tools_from_server(mock_server) # Verify tool is prefixed with server_id assert len(tools) == 1 @@ -1989,12 +1989,12 @@ async def test_get_tools_for_single_server(): with patch( "litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager" ) as mock_manager: - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) result = await _get_tools_for_single_server(mock_server, "Bearer test_token") # Verify the manager was called with correct parameters - mock_manager._get_tools_from_server.assert_called_once_with( + mock_manager.get_tools_from_server.assert_called_once_with( server=mock_server, mcp_auth_header="Bearer test_token", extra_headers=None, @@ -2050,7 +2050,7 @@ async def test_get_tools_for_single_server_applies_disallowed_tools_without_allo with patch( "litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager" ) as mock_manager: - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) result = await _get_tools_for_single_server(mock_server, "Bearer test_token") @@ -2104,7 +2104,7 @@ async def test_rest_listing_hides_key_grants_dispatch_would_refuse(): "get_allowed_tools_for_server", AsyncMock(return_value=[f"{server_id}-read_wiki_contents"]), ): - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) mock_server_manager.get_mcp_server_by_id.return_value = mock_server result = await _get_tools_for_single_server( @@ -2135,10 +2135,10 @@ async def test_list_tool_rest_api_with_server_specific_auth(): # Mock the MCPRequestHandler methods with patch.object( - MCPRequestHandler, "_get_mcp_auth_header_from_headers" + MCPRequestHandler, "get_mcp_auth_header_from_headers" ) as mock_get_auth: with patch.object( - MCPRequestHandler, "_get_mcp_server_auth_headers_from_headers" + MCPRequestHandler, "get_mcp_server_auth_headers_from_headers" ) as mock_get_server_auth: mock_get_auth.return_value = "Bearer default_token" mock_get_server_auth.return_value = { @@ -2232,10 +2232,10 @@ async def test_list_tool_rest_api_with_default_auth(): # Mock the MCPRequestHandler methods with patch.object( - MCPRequestHandler, "_get_mcp_auth_header_from_headers" + MCPRequestHandler, "get_mcp_auth_header_from_headers" ) as mock_get_auth: with patch.object( - MCPRequestHandler, "_get_mcp_server_auth_headers_from_headers" + MCPRequestHandler, "get_mcp_server_auth_headers_from_headers" ) as mock_get_server_auth: mock_get_auth.return_value = "Bearer default_token" mock_get_server_auth.return_value = {} # No server-specific headers @@ -2327,10 +2327,10 @@ async def test_list_tool_rest_api_all_servers_with_auth(): # Mock the MCPRequestHandler methods with patch.object( - MCPRequestHandler, "_get_mcp_auth_header_from_headers" + MCPRequestHandler, "get_mcp_auth_header_from_headers" ) as mock_get_auth: with patch.object( - MCPRequestHandler, "_get_mcp_server_auth_headers_from_headers" + MCPRequestHandler, "get_mcp_server_auth_headers_from_headers" ) as mock_get_server_auth: mock_get_auth.return_value = "Bearer default_token" mock_get_server_auth.return_value = { @@ -2508,7 +2508,7 @@ async def test_filter_tools_by_allowed_tools_integration(): ) # Mock the _get_tools_from_server method to return all tools - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) # Mock the MCPClient constructor with patch( @@ -2547,7 +2547,7 @@ async def test_filter_tools_by_allowed_tools_integration(): # Note: get_mcp_server_by_id is now called for each server ID instead of batch # Verify it was called with the correct server ID assert mock_manager.get_mcp_server_by_id.call_count > 0 - mock_manager._get_tools_from_server.assert_called_once() + mock_manager.get_tools_from_server.assert_called_once() @pytest.mark.asyncio @@ -2622,7 +2622,7 @@ async def test_filter_tools_by_disallowed_tools_integration(): side_effect=lambda server_ids, client_ip: (server_ids, 0) ) # Mock the _get_tools_from_server method to return all tools - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) # Mock the MCPClient constructor with patch( @@ -2661,7 +2661,7 @@ async def test_filter_tools_by_disallowed_tools_integration(): # Note: get_mcp_server_by_id is now called for each server ID instead of batch # Verify it was called with the correct server ID assert mock_manager.get_mcp_server_by_id.call_count > 0 - mock_manager._get_tools_from_server.assert_called_once() + mock_manager.get_tools_from_server.assert_called_once() @pytest.mark.asyncio @@ -2724,7 +2724,7 @@ async def test_filter_tools_no_restrictions_integration(): ) # Mock the _get_tools_from_server method to return all tools - mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) + mock_manager.get_tools_from_server = AsyncMock(return_value=mock_tools) # Mock the MCPClient constructor with patch( @@ -2988,7 +2988,7 @@ async def test_call_mcp_tool_uses_manager_permission_lookup(): ), patch.object( global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=mock_server, ) as mock_get_server, patch( @@ -3064,7 +3064,7 @@ async def test_call_mcp_tool_resolves_unprefixed_tool_name_and_checks_permission ), patch.object( global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=mock_server, ) as mock_get_server, patch( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 9aaba0e9356..5bd4ff5d894 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -54,8 +54,8 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( _deserialize_json_list, _normalize_mcp_server_cost_info, _obo_retry_applies, - _resolve_openapi_tool_auth, - _should_strip_caller_authorization, + resolve_openapi_tool_auth, + should_strip_caller_authorization, listed_tools_caller_for, ) from litellm.proxy._types import ( @@ -1366,7 +1366,7 @@ class TestMCPServerManager: manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) - with patch.object(manager, "_get_tools_from_server", new=AsyncMock()) as get_tools: + with patch.object(manager, "get_tools_from_server", new=AsyncMock()) as get_tools: await manager._initialize_tool_name_to_mcp_server_name_mapping() get_tools.assert_not_awaited() @@ -1883,7 +1883,7 @@ class TestMCPServerManager: tool1.name = "zapier_tool_1" return [tool1] - manager._get_tools_from_server = mock_get_tools_from_server + manager.get_tools_from_server = mock_get_tools_from_server # Test with server-specific auth headers mcp_server_auth_headers = { @@ -1928,7 +1928,7 @@ class TestMCPServerManager: tool.name = "github_tool_1" return [tool] - manager._get_tools_from_server = mock_get_tools_from_server + manager.get_tools_from_server = mock_get_tools_from_server # Test with only legacy auth header (no server-specific headers) result = await manager.list_tools( @@ -1965,7 +1965,7 @@ class TestMCPServerManager: tool.name = "github_tool_1" return [tool] - manager._get_tools_from_server = mock_get_tools_from_server + manager.get_tools_from_server = mock_get_tools_from_server # Test with both legacy and server-specific headers result = await manager.list_tools( @@ -2001,7 +2001,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) result = await manager._call_regular_mcp_tool( mcp_server=server, @@ -2029,7 +2029,7 @@ class TestMCPServerManager: captured["subject_token"] = subject_token return AsyncMock() - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) manager._fetch_tools_with_timeout = AsyncMock(return_value=[]) await manager._get_tools_from_server(server=server, oauth2_headers=oauth2_headers, raw_headers=raw_headers) return captured["subject_token"] @@ -2102,7 +2102,7 @@ class TestMCPServerManager: challenge = ( 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/te-401-server", error="invalid_token"' ) - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException(status_code=401, detail="Unauthorized", headers={"WWW-Authenticate": challenge}) ) with pytest.raises(MCPUpstreamAuthError) as exc_info: @@ -2128,7 +2128,7 @@ class TestMCPServerManager: client_secret="csec", ) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException(status_code=412, detail="token exchange endpoint is not configured") ) with pytest.raises(MCPServerListError) as exc_info: @@ -2178,7 +2178,7 @@ class TestMCPServerManager: manager = MCPServerManager() mock_client = AsyncMock() mock_client.call_tool = AsyncMock(side_effect=self._upstream_status_error(401, challenge)) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) with pytest.raises(MCPUpstreamAuthError) as exc_info: await self._run_call_regular(manager, server) @@ -2198,7 +2198,7 @@ class TestMCPServerManager: expected = CallToolResult(content=[], isError=is_error) mock_client = AsyncMock() mock_client.call_tool = AsyncMock(return_value=expected) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) result = await self._run_call_regular(manager, server) @@ -2219,7 +2219,7 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.call_tool = AsyncMock(side_effect=self._upstream_status_error(status_code)) mock_client.error_tool_result = MCPClient.error_tool_result - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) import litellm.proxy._experimental.mcp_server.mcp_server_manager as _mgr_mod @@ -2246,7 +2246,7 @@ class TestMCPServerManager: manager = MCPServerManager() mock_client = AsyncMock() mock_client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False)) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) result = await manager._call_regular_mcp_tool( mcp_server=server, @@ -2818,7 +2818,7 @@ class TestMCPServerManager: captured["subject_token"] = subject_token return AsyncMock() - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await call(manager) return captured.get("subject_token") @@ -3431,7 +3431,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, @@ -3472,7 +3472,7 @@ class TestMCPServerManager: ) # Migrated authorization_code => the centralized strip decision says drop the # caller's Authorization (the v2 resolver injects the stored token). - assert _should_strip_caller_authorization(mcp_server=server, raw_headers=None, user_api_key_auth=None) is True + assert should_strip_caller_authorization(mcp_server=server, raw_headers=None, user_api_key_auth=None) is True mock_client = AsyncMock() mock_client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False)) @@ -3490,7 +3490,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, @@ -3558,7 +3558,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, @@ -3615,7 +3615,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, @@ -3644,7 +3644,7 @@ class TestMCPServerManager: captured["extra_headers"] = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await manager._call_regular_mcp_tool( mcp_server=server, original_tool_name="tool", @@ -3796,7 +3796,7 @@ class TestMCPServerManager: auth_type=MCPAuth.true_passthrough, ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=true_passthrough, raw_headers={"authorization": "Bearer upstream"}, user_api_key_auth=UserAPIKeyAuth(api_key=None), @@ -3812,7 +3812,7 @@ class TestMCPServerManager: auth_type=MCPAuth.oauth_delegate, ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=oauth_delegate, raw_headers={ "x-litellm-api-key": "Bearer sk-litellm-key", @@ -3823,7 +3823,7 @@ class TestMCPServerManager: is False ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=oauth_delegate, raw_headers={"authorization": "Bearer sk-litellm-key"}, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), @@ -3845,7 +3845,7 @@ class TestMCPServerManager: auth_type=MCPAuth.oauth_delegate, ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=oauth_delegate, raw_headers={"authorization": "Bearer eyJ-idp-jwt"}, user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key=None), @@ -3853,7 +3853,7 @@ class TestMCPServerManager: is True ) assert ( - _should_strip_caller_authorization( + should_strip_caller_authorization( mcp_server=oauth_delegate, raw_headers={ "x-litellm-api-key": "Bearer sk-9876", @@ -4055,7 +4055,7 @@ class TestMCPServerManager: def test_caller_authorization_fans_out_only_with_second_consumer(self): from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _caller_authorization_fans_out, + caller_authorization_fans_out, ) delegate = MCPServer( @@ -4081,10 +4081,10 @@ class TestMCPServerManager: authentication_token="x", ) - assert _caller_authorization_fans_out(delegate, None) is False - assert _caller_authorization_fans_out(delegate, [delegate]) is False - assert _caller_authorization_fans_out(delegate, [delegate, static_server]) is False - assert _caller_authorization_fans_out(delegate, [delegate, second]) is True + assert caller_authorization_fans_out(delegate, None) is False + assert caller_authorization_fans_out(delegate, [delegate]) is False + assert caller_authorization_fans_out(delegate, [delegate, static_server]) is False + assert caller_authorization_fans_out(delegate, [delegate, second]) is True @pytest.mark.asyncio async def test_get_prompts_from_server_success(self): @@ -4107,7 +4107,7 @@ class TestMCPServerManager: with patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ): @@ -4140,7 +4140,7 @@ class TestMCPServerManager: with patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ): @@ -4181,7 +4181,7 @@ class TestMCPServerManager: with ( patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ) as mock_create_client, @@ -4236,7 +4236,7 @@ class TestMCPServerManager: with ( patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ) as mock_create_client, @@ -4291,7 +4291,7 @@ class TestMCPServerManager: with patch.object( manager, - "_create_mcp_client", + "create_mcp_client", new_callable=AsyncMock, return_value=mock_client, ) as mock_create_client: @@ -4854,7 +4854,7 @@ class TestMCPServerManager: tool.name = "github_tool_1" return [tool] - manager._get_tools_from_server = mock_get_tools_from_server + manager.get_tools_from_server = mock_get_tools_from_server # Test with server-specific headers that match server_name (even without alias) result = await manager.list_tools( @@ -5011,7 +5011,7 @@ class TestMCPServerManager: # Mock successful client.run_with_session mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(return_value="ok") - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("test-server") @@ -5043,7 +5043,7 @@ class TestMCPServerManager: # Mock failed client.run_with_session mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(side_effect=Exception("Connection timeout")) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("test-server") @@ -5069,7 +5069,7 @@ class TestMCPServerManager: ) manager.get_mcp_server_by_id = MagicMock(return_value=server) manager._resolve_static_headers_with_env_vars = AsyncMock(return_value=None) - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException(status_code=503, detail="OAuth discovery unavailable") ) @@ -5119,12 +5119,12 @@ class TestMCPServerManager: static_headers={"Authorization": "Bearer static-secret", "X-API-Key": "key-secret", "Cookie": "secret"}, ) manager.registry[server.server_id] = server - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() route: Final = respx_mock.get(server.url).respond(401) result: Final = await manager.health_check_server(server.server_id, mcp_auth_header="caller-secret") - manager._create_mcp_client.assert_not_called() + manager.create_mcp_client.assert_not_called() assert result.status == "reachable" assert result.health_check_error is None assert result.last_health_check is not None @@ -5167,12 +5167,12 @@ class TestMCPServerManager: url="http://no-token-server.com", ) manager.registry[server.server_id] = server - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() route: Final = respx_mock.get(server.url).respond(response_code) result: Final = await manager.health_check_server(server.server_id) - manager._create_mcp_client.assert_not_called() + manager.create_mcp_client.assert_not_called() assert route.call_count == 1 assert result.status == "reachable" assert result.health_check_error is None @@ -5435,7 +5435,7 @@ class TestMCPServerManager: captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + manager.create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) # Perform health check result = await manager.health_check_server("test-server") @@ -5465,12 +5465,12 @@ class TestMCPServerManager: extra_headers=["Authorization"], ) manager.registry[server.server_id] = server - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() route: Final = respx_mock.get(server.url).respond(401) result: Final = await manager.health_check_server(server.server_id) - manager._create_mcp_client.assert_not_called() + manager.create_mcp_client.assert_not_called() assert route.call_count == 1 assert "authorization" not in route.calls[0].request.headers assert result.status == "reachable" @@ -5493,12 +5493,12 @@ class TestMCPServerManager: extra_headers=["x-api-key"], ) manager.registry[server.server_id] = server - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() route: Final = respx_mock.get(server.url).respond(403) result: Final = await manager.health_check_server(server.server_id) - manager._create_mcp_client.assert_not_called() + manager.create_mcp_client.assert_not_called() assert route.call_count == 1 assert "x-api-key" not in route.calls[0].request.headers assert result.status == "reachable" @@ -5526,13 +5526,13 @@ class TestMCPServerManager: # Mock successful client mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(return_value="ok") - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("public-server") # Verify that client WAS created (health check should run) - manager._create_mcp_client.assert_called_once() + manager.create_mcp_client.assert_called_once() # Verify results assert isinstance(result, LiteLLM_MCPServerTable) @@ -5562,13 +5562,13 @@ class TestMCPServerManager: # Mock successful client mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(return_value="ok") - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("custom-server") # Verify that client WAS created (health check should run) - manager._create_mcp_client.assert_called_once() + manager.create_mcp_client.assert_called_once() # Verify results assert isinstance(result, LiteLLM_MCPServerTable) @@ -5858,8 +5858,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # This should not raise an exception @@ -5925,8 +5925,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # This should not raise an exception @@ -5992,8 +5992,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # This should not raise an exception @@ -6027,8 +6027,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # tool2 should be allowed since it's in allowed_tools (takes precedence) @@ -6067,7 +6067,7 @@ class TestMCPServerManager: ) # Mock client creation and fetching tools - manager._create_mcp_client = AsyncMock(return_value=object()) + manager.create_mcp_client = AsyncMock(return_value=object()) # Tools returned upstream (unprefixed from provider) upstream_tool = MCPTool( @@ -6105,7 +6105,7 @@ class TestMCPServerManager: transport=MCPTransport.http, ) - manager._create_mcp_client = AsyncMock(return_value=object()) + manager.create_mcp_client = AsyncMock(return_value=object()) manager._fetch_tools_with_timeout = AsyncMock(return_value=[]) user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="alice") @@ -6739,7 +6739,7 @@ class TestMCPServerManager: with patch.object( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", new=AsyncMock(return_value=[tool1, tool2, tool3]), ): # Call the REST endpoint helper @@ -6780,7 +6780,7 @@ class TestMCPServerManager: with patch.object( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", new=AsyncMock(return_value=[tool1, tool2, tool3]), ): # Call the REST endpoint helper @@ -6819,7 +6819,7 @@ class TestMCPServerManager: with patch.object( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", new=AsyncMock(return_value=[tool1, tool2]), ): # Call the REST endpoint helper @@ -6893,8 +6893,8 @@ class TestMCPServerManager: ) proxy_logging = _mock_proxy_logging() - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) # Should succeed @@ -6936,8 +6936,8 @@ class TestMCPServerManager: ) proxy_logging = _mock_proxy_logging() - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) # Should fail with 403 @@ -7079,8 +7079,8 @@ class TestMCPServerManager: proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) # Test 1: Call getpetbyid (unprefixed in allowed_tools) - should succeed @@ -7156,14 +7156,14 @@ class TestMCPServerManager: mock_client.call_tool.side_effect = mock_call_tool # Mock _create_mcp_client to return our mock client - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) user_api_key_auth: Final = UserAPIKeyAuth(api_key="sk-test") # Mock proxy logging proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) @@ -7205,11 +7205,11 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) return manager, proxy_logging_obj @@ -7235,7 +7235,7 @@ class TestMCPServerManager: proxy_logging_obj=proxy_logging_obj, ) - hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging_obj.create_mcp_request_object_from_kwargs.call_args.args[0] assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ("Runs the test tool", schema) @pytest.mark.asyncio @@ -7279,7 +7279,7 @@ class TestMCPServerManager: proxy_logging_obj=proxy_logging_obj, ) - hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging_obj.create_mcp_request_object_from_kwargs.call_args.args[0] assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None) def test_get_listed_tool_resolves_the_bare_name_from_the_latest_listing(self): @@ -7384,7 +7384,7 @@ class TestMCPServerManager: await release_fetch.wait() return [MCPTool(name="turn", description="before save", inputSchema={})] - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) manager._fetch_tools_with_timeout = fetch caller = ListedToolsCaller(user_api_key_auth=user) @@ -7636,7 +7636,7 @@ class TestMCPServerManager: auth_type=MCPAuth.api_key, ) user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-cold-user") - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="turn", description="listed while db down", inputSchema={})] ) @@ -7688,13 +7688,13 @@ class TestMCPServerManager: user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-user") mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="turn", description="stored cred catalog", inputSchema={})] ) proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) cache_byok_credential("byok-user", "byok-catalog", "stored-secret") @@ -7718,7 +7718,7 @@ class TestMCPServerManager: finally: byok_credential_cache.delete_cache(byok_credential_cache_key("byok-user", "byok-catalog")) - hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging_obj.create_mcp_request_object_from_kwargs.call_args.args[0] assert hook_kwargs["tool_description"] == "stored cred catalog" @pytest.mark.asyncio @@ -7731,7 +7731,7 @@ class TestMCPServerManager: url="http://byok-catalog", is_byok=True, ) - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="turn", description="t", inputSchema={})] ) @@ -7749,7 +7749,7 @@ class TestMCPServerManager: ) listed = manager.get_listed_tool(server, "turn", caller) assert listed is not None and listed.description == "t" - assert manager._create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" + assert manager.create_mcp_client.await_args.kwargs["mcp_auth_header"] == "Bearer hdr" @pytest.mark.parametrize( "server_auth", @@ -7794,7 +7794,7 @@ class TestMCPServerManager: **server_auth, ) alice = UserAPIKeyAuth(api_key="sk-alice", user_id="alice") - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="echo", description="listed catalog", inputSchema={})] ) @@ -7815,7 +7815,7 @@ class TestMCPServerManager: finally: byok_credential_cache.delete_cache(byok_credential_cache_key("alice", "cc1")) - client_kwargs = manager._create_mcp_client.await_args.kwargs + client_kwargs = manager.create_mcp_client.await_args.kwargs assert client_kwargs["mcp_auth_header"] is None, client_kwargs assert client_kwargs["extra_headers"] == {"Authorization": "Bearer signed-jwt"} signer_headers.assert_awaited_once() @@ -8014,10 +8014,10 @@ class TestMCPServerManager: } mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) - manager._create_mcp_client = AsyncMock(return_value=mock_client) + manager.create_mcp_client = AsyncMock(return_value=mock_client) manager._fetch_tools_with_timeout = AsyncMock(side_effect=lambda client, name: catalogs[client.workspace]) for workspace in ("A", "B"): - manager._create_mcp_client.return_value.workspace = workspace + manager.create_mcp_client.return_value.workspace = workspace await manager._get_tools_from_server( server=server, extra_headers={"X-Workspace": workspace}, @@ -8027,8 +8027,8 @@ class TestMCPServerManager: ) proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) await manager.call_tool( @@ -8040,7 +8040,7 @@ class TestMCPServerManager: raw_headers={"x-workspace": "A", "authorization": "Bearer sk-litellm"}, ) - hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging_obj.create_mcp_request_object_from_kwargs.call_args.args[0] assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ( "Catalog A", {"properties": {"turn": {"description": "A"}}}, @@ -8093,7 +8093,7 @@ class TestMCPServerManager: spec_path="/spec.yaml", ) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) async def _handler(**kwargs): return None @@ -8128,7 +8128,7 @@ class TestMCPServerManager: spec_path="/spec.yaml", ) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) async def _handler(**kwargs): return None @@ -8170,7 +8170,7 @@ class TestMCPServerManager: server_id="srv", name="srv", alias="srv", transport=MCPTransport.http, url=None, spec_path="/spec.yaml" ) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) global_mcp_tool_registry.unregister_tools_with_prefix("srv-") global_mcp_tool_registry.register_tool( name="srv-echo", description="Echoes", input_schema={"type": "object"}, handler=lambda **kwargs: None @@ -8433,7 +8433,7 @@ class TestMCPServerManager: MCPServerAccess, ) from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_active_toolset_id, + mcp_active_toolset_id, ) from litellm.proxy._types import UserAPIKeyAuth @@ -8456,7 +8456,7 @@ class TestMCPServerManager: user_api_key_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-123") user_api_key_auth.mcp_toolset_id = "toolset-abc" - token = _mcp_active_toolset_id.set("unrelated-ambient-toolset") + token = mcp_active_toolset_id.set("unrelated-ambient-toolset") try: with ( patch.object(proxy_server_module, "user_api_key_cache", cache), @@ -8474,7 +8474,7 @@ class TestMCPServerManager: ): result = await manager.get_allowed_mcp_servers(user_api_key_auth) finally: - _mcp_active_toolset_id.reset(token) + mcp_active_toolset_id.reset(token) assert result == ["toolset-server"] @@ -9599,7 +9599,7 @@ class TestMCPServerTimestamps: timeout=0.01, ) - with patch.object(manager, "_create_mcp_client", return_value=mock_client): + with patch.object(manager, "create_mcp_client", return_value=mock_client): with pytest.raises(HTTPException) as exc_info: await manager._call_regular_mcp_tool( mcp_server=server, @@ -10770,7 +10770,7 @@ class TestHealthCheckInterpolatesGlobalEnvVars: captured["extra_headers"] = extra_headers return mock_client - manager._create_mcp_client = AsyncMock(side_effect=_create) + manager.create_mcp_client = AsyncMock(side_effect=_create) return captured @pytest.mark.asyncio @@ -11331,7 +11331,7 @@ class TestMCPToolsListAuthSurfacing: manager = MCPServerManager() server = MCPServer(server_id="oauth-srv", name="oauth-srv", transport=MCPTransport.http) challenge = 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/oauth-srv"' - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException( status_code=401, detail="Unauthorized", @@ -11355,7 +11355,7 @@ class TestMCPToolsListAuthSurfacing: 401/403 remain the challenge-class statuses routed to MCPUpstreamAuthError.""" manager = MCPServerManager() server = MCPServer(server_id="stdio-srv", name="stdio-srv", transport=MCPTransport.http) - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException( status_code=500, detail="MCP stdio command 'foo' is not in the allowlist", @@ -11393,7 +11393,7 @@ class TestMCPToolsListAuthSurfacing: ) wrapper = RuntimeError("client build failed") wrapper.__cause__ = causal - manager._create_mcp_client = AsyncMock(side_effect=wrapper) + manager.create_mcp_client = AsyncMock(side_effect=wrapper) with pytest.raises(MCPUpstreamAuthError) as exc_info: await manager._get_tools_from_server(server) @@ -11432,7 +11432,7 @@ class TestMCPToolsListAuthSurfacing: ) wrapper = RuntimeError("client build failed") wrapper.__cause__ = causal - manager._create_mcp_client = AsyncMock(side_effect=wrapper) + manager.create_mcp_client = AsyncMock(side_effect=wrapper) with pytest.raises(MCPUpstreamAuthError) as exc_info: await manager._get_tools_from_server(bridge_server) @@ -11462,7 +11462,7 @@ class TestMCPToolsListAuthSurfacing: upstream_challenge = 'Bearer resource_metadata="https://upstream.example/.well-known/oauth-protected-resource"' client = MagicMock() client.list_tools = AsyncMock(side_effect=_upstream_status_error(401, upstream_challenge)) - manager._create_mcp_client = AsyncMock(return_value=client) + manager.create_mcp_client = AsyncMock(return_value=client) with pytest.raises(MCPUpstreamAuthError) as exc_info: await manager._get_tools_from_server(bridge_server) @@ -11487,7 +11487,7 @@ class TestMCPToolsListAuthSurfacing: auth_type=MCPAuth.oauth_delegate, dcr_bridge=True, ) - manager._create_mcp_client = AsyncMock( + manager.create_mcp_client = AsyncMock( side_effect=HTTPException( status_code=401, detail="Unauthorized", @@ -11525,7 +11525,7 @@ class TestMCPToolsListAuthSurfacing: ) return [good_tool] - manager._get_tools_from_server = fake_get_tools + manager.get_tools_from_server = fake_get_tools result = await manager.list_tools() @@ -11544,7 +11544,7 @@ def test_should_strip_caller_authorization_for_token_exchange(): client_id="cid", client_secret="csec", ) - assert _should_strip_caller_authorization(mcp_server=server, raw_headers=None, user_api_key_auth=None) is True + assert should_strip_caller_authorization(mcp_server=server, raw_headers=None, user_api_key_auth=None) is True def _retry_gate_server(auth_type: MCPAuthType) -> MCPServer: @@ -11632,7 +11632,7 @@ class TestOBOCallToolRetry: success = CallToolResult(content=[], isError=False) first = _RetryFakeClient(raises=_UpstreamAuthError(401)) retry = _RetryFakeClient(result=success) - manager._create_mcp_client = AsyncMock(return_value=retry) + manager.create_mcp_client = AsyncMock(return_value=retry) result = await manager._obo_call_tool_with_retry( client=first, @@ -11648,7 +11648,7 @@ class TestOBOCallToolRetry: assert result is success manager._cred_provider.invalidate_credentials.assert_awaited_once() - manager._create_mcp_client.assert_awaited_once() + manager.create_mcp_client.assert_awaited_once() assert first.attempts == 1 and retry.attempts == 1 @pytest.mark.asyncio @@ -11663,7 +11663,7 @@ class TestOBOCallToolRetry: success = CallToolResult(content=[], isError=False) first = _RetryFakeClient(raises=_UpstreamAuthError(401)) retry = _RetryFakeClient(result=success) - manager._create_mcp_client = AsyncMock(return_value=retry) + manager.create_mcp_client = AsyncMock(return_value=retry) server = MCPServer( server_id="id-jag-srv", name="id-jag", @@ -11702,7 +11702,7 @@ class TestOBOCallToolRetry: success = CallToolResult(content=[], isError=False) first = _RetryFakeClient(raises=_UpstreamAuthError(401)) retry = _RetryFakeClient(result=success) - manager._create_mcp_client = AsyncMock(side_effect=[first, retry]) + manager.create_mcp_client = AsyncMock(side_effect=[first, retry]) server = MCPServer( server_id="id-jag-srv", name="id-jag", @@ -11735,7 +11735,7 @@ class TestOBOCallToolRetry: async def test_non_auth_error_does_not_retry(self): manager = self._manager() first = _RetryFakeClient(raises=ValueError("tool blew up")) - manager._create_mcp_client = AsyncMock() + manager.create_mcp_client = AsyncMock() result = await manager._obo_call_tool_with_retry( client=first, @@ -11751,7 +11751,7 @@ class TestOBOCallToolRetry: assert result.is_error is True manager._cred_provider.invalidate_credentials.assert_not_awaited() - manager._create_mcp_client.assert_not_awaited() + manager.create_mcp_client.assert_not_awaited() assert first.attempts == 1 @pytest.mark.asyncio @@ -11760,7 +11760,7 @@ class TestOBOCallToolRetry: first = _RetryFakeClient(raises=_UpstreamAuthError(401)) # The retry client still fails; with raise_on_error defaulting False it returns isError. retry = _RetryFakeClient(raises=_UpstreamAuthError(401)) - manager._create_mcp_client = AsyncMock(return_value=retry) + manager.create_mcp_client = AsyncMock(return_value=retry) result = await manager._obo_call_tool_with_retry( client=first, @@ -11775,7 +11775,7 @@ class TestOBOCallToolRetry: ) assert result.is_error is True - manager._create_mcp_client.assert_awaited_once() + manager.create_mcp_client.assert_awaited_once() assert first.attempts == 1 and retry.attempts == 1 @@ -11819,7 +11819,7 @@ class TestOBOConcurrencyLimit: return CallToolResult(content=[], isError=False) manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=_ConcurrencyRecordingClient()) + manager.create_mcp_client = AsyncMock(return_value=_ConcurrencyRecordingClient()) async def _dispatch(): return await manager._call_regular_mcp_tool( @@ -12041,7 +12041,7 @@ async def test_aggregate_list_still_absorbs_step_up_challenged_server(): ) return [good_tool] - manager._get_tools_from_server = fake_get_tools + manager.get_tools_from_server = fake_get_tools result = await manager.list_tools() @@ -12869,8 +12869,8 @@ def _mock_proxy_logging() -> MagicMock: def _permissive_proxy_logging() -> MagicMock: proxy_logging_obj: Final = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) return proxy_logging_obj @@ -13733,7 +13733,7 @@ class TestResolveOpenapiToolAuth: expected_extra_keys: set, expected_credential: object, ): - auth_value, forwarded, credential = _resolve_openapi_tool_auth( + auth_value, forwarded, credential = resolve_openapi_tool_auth( mcp_server=self._server(), mcp_auth_header=byok, mcp_server_auth_headers=per_server, @@ -13747,7 +13747,7 @@ class TestResolveOpenapiToolAuth: def test_per_server_value_is_never_re_prefixed(self): """The regression that a naive wiring produces: the caller already sent ``Bearer ``.""" - auth_value, _, credential = _resolve_openapi_tool_auth( + auth_value, _, credential = resolve_openapi_tool_auth( mcp_server=self._server(auth_type=MCPAuth.api_key), mcp_auth_header="byok-secret", mcp_server_auth_headers={"report_api": "Bearer caller-token"}, @@ -13762,7 +13762,7 @@ class TestResolveOpenapiToolAuth: def test_per_server_authorization_is_not_also_left_in_forwarded_headers(self): """``resolve_openapi_upstream_auth`` pops Authorization out of the forwarded headers, so a second copy there would give the passthrough arm two sources to reconcile.""" - _, forwarded, _ = _resolve_openapi_tool_auth( + _, forwarded, _ = resolve_openapi_tool_auth( mcp_server=self._server(), mcp_auth_header=None, mcp_server_auth_headers={"report_api": {"Authorization": "Bearer caller-token"}}, @@ -14518,12 +14518,12 @@ class TestLitellmAdmissionKeyIsNeverTheSubjectToken: client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False)) client.list_prompts_result = AsyncMock(return_value=ListPromptsResult(prompts=[])) client.read_resource = AsyncMock(return_value=ReadResourceResult(contents=[])) - manager._create_mcp_client = AsyncMock(return_value=client) + manager.create_mcp_client = AsyncMock(return_value=client) return manager @staticmethod def _subject_token_given_to_client(manager: MCPServerManager) -> str | None: - return manager._create_mcp_client.call_args.kwargs["subject_token"] + return manager.create_mcp_client.call_args.kwargs["subject_token"] async def _call_tool_subject(self, server: MCPServer, oauth2_headers, raw_headers, user_api_key_auth): manager: Final = self._manager_with_recording_client() @@ -15549,7 +15549,7 @@ async def test_openapi_listing_ignores_overlapping_server_prefix() -> None: from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry manager: Final = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) for prefix in ("pet-", "petstore-"): global_mcp_tool_registry.unregister_tools_with_prefix(prefix) _register_local_tool("pet-list", "Local pet tool") @@ -15570,7 +15570,7 @@ async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefi from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry manager: Final = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager.create_mcp_client = AsyncMock(return_value=AsyncMock()) global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") _register_local_tool("pet_store-list", "Pet store tool") try: @@ -16151,9 +16151,9 @@ class TestProtectedCredentialPreparation: caller: str | None, ) -> None: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, create_tool_function, + request_auth_header, + request_extra_headers, ) tool: Final = create_tool_function( @@ -16166,8 +16166,8 @@ class TestProtectedCredentialPreparation: ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") - caller_token: Final = _request_auth_header.set(caller) - extra_token: Final = _request_extra_headers.set(forwarded) + caller_token: Final = request_auth_header.set(caller) + extra_token: Final = request_extra_headers.set(forwarded) try: assert await tool() == TextResult("authenticated") sent: Final = destination.calls.last.request.headers @@ -16176,8 +16176,8 @@ class TestProtectedCredentialPreparation: assert sent["authorization"] == caller assert destination.call_count == 1 finally: - _request_auth_header.reset(caller_token) - _request_extra_headers.reset(extra_token) + request_auth_header.reset(caller_token) + request_extra_headers.reset(extra_token) @pytest.mark.asyncio async def test_static_resolution_cancellation_closes_flow(self) -> None: @@ -17925,7 +17925,7 @@ def catalog_guardrail(monkeypatch): def _catalog_manager(*upstream_tools: MCPTool) -> MCPServerManager: manager = MCPServerManager() - manager._create_mcp_client = AsyncMock(return_value=object()) + manager.create_mcp_client = AsyncMock(return_value=object()) manager._fetch_tools_with_timeout = AsyncMock(return_value=list(upstream_tools)) return manager @@ -18436,8 +18436,8 @@ class TestToolCatalogGuard: server = _notes_server({"list_notes": _pin(LIST_NOTES)}) user_api_key_auth = MagicMock(object_permission=None, object_permission_id=None) proxy_logging_obj = _mock_proxy_logging() - proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) - proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj.convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) with pytest.raises(HTTPException) as exc_info: @@ -18910,7 +18910,7 @@ async def test_catalog_page_registers_bare_routes_only_for_complete_initial_disc other = MCPServer(server_id="other", name="other", transport=MCPTransport.http) manager.registry = {server.server_id: server, other.server_id: other} manager._create_prefixed_tools([LIST_NOTES], other) - manager._create_mcp_client.return_value = SimpleNamespace( + manager.create_mcp_client.return_value = SimpleNamespace( list_tools_page=AsyncMock(return_value=ListToolsResult(tools=[LIST_NOTES], next_cursor=next_cursor)) ) @@ -18943,7 +18943,7 @@ async def test_paginated_listing_keeps_earlier_tool_metadata_and_caller_isolatio ListToolsResult(tools=[second]), ListToolsResult(tools=[second]), ] - manager._create_mcp_client = AsyncMock(return_value=client) + manager.create_mcp_client = AsyncMock(return_value=client) monkeypatch.setattr(operations, "global_mcp_server_manager", manager) caller = UserAPIKeyAuth(api_key="owned-caller", user_id="alice") context = operations.prepare_context(caller) @@ -19025,7 +19025,7 @@ async def test_failed_aggregate_continuation_preserves_only_delivered_tool_metad ] async def create_client(server, **kwargs): return clients[server.server_id] - manager._create_mcp_client = create_client + manager.create_mcp_client = create_client monkeypatch.setattr(operations, "global_mcp_server_manager", manager) caller = UserAPIKeyAuth(api_key="owned-caller", user_id="alice") context = operations.prepare_context(caller) @@ -19066,7 +19066,7 @@ async def test_aggregate_publishes_complete_bare_routes_only_after_delivering_a_ async def create_client(server, **kwargs): return clients[server.server_id] - manager._create_mcp_client = create_client + manager.create_mcp_client = create_client monkeypatch.setattr(operations, "global_mcp_server_manager", manager) context = operations.prepare_context(UserAPIKeyAuth(api_key="owned-caller", user_id="alice")) listing = catalog.aggregate_gateway_tools(context, PaginatedRequestParams(), servers, {}, record_listing=True) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index df3f1d78a7e..97c6e2703ad 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -826,7 +826,7 @@ async def test_call_tool_m2m_skips_authorization_headers(): mock_client = MagicMock() mock_client.call_tool = AsyncMock(return_value=MagicMock()) - with patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=mock_client)) as create_client_mock: + with patch.object(manager, "create_mcp_client", new=AsyncMock(return_value=mock_client)) as create_client_mock: await manager._call_regular_mcp_tool( mcp_server=server, original_tool_name="echo", @@ -1302,7 +1302,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails(): # Failing server raises an exception raise Exception("Server connection failed") - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -1398,7 +1398,7 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing(): # All servers fail raise Exception(f"Server {server.name} connection failed") - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -4237,11 +4237,11 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name(): mock_get_allowed, ), patch( - "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.get_mcp_servers_from_access_groups", mock_db_lookup, ), patch( - "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager._get_tools_from_server", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_tools_from_server", mock_get_tools_spy, ), ): @@ -4341,7 +4341,7 @@ async def test_oauth2_caller_headers_not_forwarded_for_migrated_server(): with ( patch.object( global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", side_effect=mock_create_mcp_client, ) as mock_create_client, patch.object( @@ -4438,7 +4438,7 @@ async def test_list_tools_single_server_unprefixed_names(): tool.input_schema = {} return [tool] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -4517,7 +4517,7 @@ async def test_list_tools_multiple_servers_prefixed_names(): tool.input_schema = {} return [tool] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -4557,7 +4557,7 @@ async def test_mcp_manager_allows_public_servers_without_permissions(): with ( patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", + "litellm.proxy.management_endpoints.common_utils.user_api_key_has_admin_view", return_value=False, ), patch( @@ -4592,7 +4592,7 @@ async def test_mcp_manager_returns_public_when_permission_lookup_fails(): with ( patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", + "litellm.proxy.management_endpoints.common_utils.user_api_key_has_admin_view", return_value=False, ), patch( @@ -4636,7 +4636,7 @@ async def test_mcp_manager_merges_public_and_restricted_servers(): with ( patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", + "litellm.proxy.management_endpoints.common_utils.user_api_key_has_admin_view", return_value=False, ), patch( @@ -4946,7 +4946,7 @@ async def test_list_tools_filters_by_key_team_permissions(): return [tool1, tool2, tool3, tool4] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -5057,7 +5057,7 @@ async def test_list_tools_with_team_tool_permissions_inheritance(): return [tool1, tool2, tool3, tool4] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -5149,7 +5149,7 @@ async def test_list_tools_with_no_tool_permissions_shows_all(): return [tool1, tool2, tool3] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -5255,7 +5255,7 @@ async def test_list_tools_strips_prefix_when_matching_permissions(): return [tool1, tool2, tool3, tool4] - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -5796,7 +5796,7 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab side_effect=_capture_function_setup, ), ): - mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[tool_1]) listing = await _get_tools_from_mcp_servers( user_api_key_auth=user_auth, @@ -5878,7 +5878,7 @@ async def test_get_tools_from_mcp_servers_returns_tools_when_success_logging_fai return_value=(dummy_logging_obj, None), ), ): - mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[tool_1]) listing = await _get_tools_from_mcp_servers( user_api_key_auth=user_auth, @@ -6195,7 +6195,7 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token(): new=AsyncMock(side_effect=lambda tools, **_: tools), ), ): - mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + mock_manager.get_tools_from_server = AsyncMock(return_value=[tool_1]) listing = await _get_tools_from_mcp_servers( user_api_key_auth=user_auth, @@ -6209,8 +6209,8 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token(): mock_prefetch.assert_awaited_once_with(user_auth) # The stored token was forwarded to the MCP transport layer as extra_headers - mock_manager._get_tools_from_server.assert_awaited_once() - call_kwargs = mock_manager._get_tools_from_server.await_args.kwargs + mock_manager.get_tools_from_server.assert_awaited_once() + call_kwargs = mock_manager.get_tools_from_server.await_args.kwargs assert call_kwargs["extra_headers"] == {"Authorization": f"Bearer {STORED_TOKEN}"} assert listing.tools == [tool_1] @@ -6351,7 +6351,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: ) server = _make_instruction_server(server_id="yaml-only", instructions="from yaml") - with patch.object(global_mcp_server_manager, "_create_mcp_client", AsyncMock()) as mock_create: + with patch.object(global_mcp_server_manager, "create_mcp_client", AsyncMock()) as mock_create: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) mock_create.assert_not_awaited() @@ -6366,7 +6366,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: server = _make_instruction_server(server_id="cached-only", instructions=None) global_mcp_server_manager._upstream_initialize_instructions_by_server_id["cached-only"] = "warm" try: - with patch.object(global_mcp_server_manager, "_create_mcp_client", AsyncMock()) as mock_create: + with patch.object(global_mcp_server_manager, "create_mcp_client", AsyncMock()) as mock_create: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) mock_create.assert_not_awaited() finally: @@ -6381,7 +6381,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: ) server = _make_instruction_server(server_id="openapi-spec", spec_path="/openapi.json", url=None) - with patch.object(global_mcp_server_manager, "_create_mcp_client", AsyncMock()) as mock_create: + with patch.object(global_mcp_server_manager, "create_mcp_client", AsyncMock()) as mock_create: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) mock_create.assert_not_awaited() @@ -6400,7 +6400,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: with patch.object( global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", AsyncMock(return_value=fake_client), ): try: @@ -6428,7 +6428,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: fake_client._last_initialize_instructions = None # upstream sent nothing create = AsyncMock(return_value=fake_client) - with patch.object(global_mcp_server_manager, "_create_mcp_client", create): + with patch.object(global_mcp_server_manager, "create_mcp_client", create): try: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) @@ -6453,7 +6453,7 @@ class TestEnsureUpstreamInitializeInstructionsCached: fake_client._last_initialize_instructions = None create = AsyncMock(return_value=fake_client) - with patch.object(global_mcp_server_manager, "_create_mcp_client", create): + with patch.object(global_mcp_server_manager, "create_mcp_client", create): try: await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(server) @@ -6485,22 +6485,22 @@ class TestGatewayCreateInitializationOptions: """When ContextVar is None, instructions are absent.""" try: from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_gateway_initialize_instructions, - _mcp_gateway_server_name, + mcp_gateway_initialize_instructions, + mcp_gateway_server_name, ) from litellm.proxy._experimental.mcp_server.server import server except ImportError: pytest.skip("MCP server not available") - instructions_token = _mcp_gateway_initialize_instructions.set(None) - server_name_token = _mcp_gateway_server_name.set(None) + instructions_token = mcp_gateway_initialize_instructions.set(None) + server_name_token = mcp_gateway_server_name.set(None) try: opts = server.create_initialization_options() assert getattr(opts, "instructions", None) is None assert opts.server_name == "litellm-mcp-server" finally: - _mcp_gateway_initialize_instructions.reset(instructions_token) - _mcp_gateway_server_name.reset(server_name_token) + mcp_gateway_initialize_instructions.reset(instructions_token) + mcp_gateway_server_name.reset(server_name_token) @pytest.mark.asyncio async def test_scoped_request_uses_configured_server_alias(self): @@ -6529,7 +6529,7 @@ class TestGatewayCreateInitializationOptions: ), patch.object( global_mcp_server_manager, - "_ensure_upstream_initialize_instructions_cached", + "ensure_upstream_initialize_instructions_cached", new_callable=AsyncMock, ), ): @@ -6599,7 +6599,7 @@ class TestGatewayCreateInitializationOptions: async def test_non_initialize_request_with_no_granted_servers_is_not_rejected_here(self): from litellm.proxy._experimental.mcp_server.server import ( _gateway_initialize_instructions_request_scope, - _mcp_gateway_initialize_instructions, + mcp_gateway_initialize_instructions, ) from litellm.proxy._types import UserAPIKeyAuth @@ -6613,7 +6613,7 @@ class TestGatewayCreateInitializationOptions: mcp_servers=None, client_ip=None, ): - assert _mcp_gateway_initialize_instructions.get() is None + assert mcp_gateway_initialize_instructions.get() is None @pytest.mark.asyncio async def test_sse_handler_scopes_server_name_from_single_server_path(self): @@ -6678,7 +6678,7 @@ class TestGatewayCreateInitializationOptions: ), patch.object( global_mcp_server_manager, - "_ensure_upstream_initialize_instructions_cached", + "ensure_upstream_initialize_instructions_cached", new_callable=AsyncMock, ), patch( @@ -6708,31 +6708,31 @@ class TestGatewayCreateInitializationOptions: """When ContextVar has a value, it appears in InitializationOptions.""" try: from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_gateway_initialize_instructions, + mcp_gateway_initialize_instructions, ) from litellm.proxy._experimental.mcp_server.server import server except ImportError: pytest.skip("MCP server not available") - tok = _mcp_gateway_initialize_instructions.set("hello from merge") + tok = mcp_gateway_initialize_instructions.set("hello from merge") try: opts = server.create_initialization_options() assert opts.instructions == "hello from merge" finally: - _mcp_gateway_initialize_instructions.reset(tok) + mcp_gateway_initialize_instructions.reset(tok) def test_contextvar_reset_removes_instructions(self): """After resetting the ContextVar, instructions disappear.""" try: from litellm.proxy._experimental.mcp_server.mcp_context import ( - _mcp_gateway_initialize_instructions, + mcp_gateway_initialize_instructions, ) from litellm.proxy._experimental.mcp_server.server import server except ImportError: pytest.skip("MCP server not available") - tok = _mcp_gateway_initialize_instructions.set("temporary") - _mcp_gateway_initialize_instructions.reset(tok) + tok = mcp_gateway_initialize_instructions.set("temporary") + mcp_gateway_initialize_instructions.reset(tok) opts = server.create_initialization_options() assert getattr(opts, "instructions", None) is None @@ -6821,7 +6821,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow(): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"]) mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server) mock_manager.filter_server_ids_by_ip_with_info = MagicMock(return_value=(["legacy-m2m-id"], 0)) - mock_manager._get_tools_from_server = AsyncMock(side_effect=capture_extra_headers) + mock_manager.get_tools_from_server = AsyncMock(side_effect=capture_extra_headers) listing = await _get_tools_from_mcp_servers( user_api_key_auth=user_auth, @@ -6891,7 +6891,7 @@ async def test_call_tool_empty_extra_headers_returns_none(): with ( patch.object( manager, - "_create_mcp_client", + "create_mcp_client", side_effect=capture_create_mcp_client, ), patch.object( @@ -7488,7 +7488,7 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=oauth_server, ), patch.object( @@ -7553,7 +7553,7 @@ def _worker_that_never_listed(server: MCPServer, upstream_tools: tuple[str, ...] with ( patch.object( # test-quality-ok: the upstream MCP session is the boundary; a real one needs an initialize handshake over a live server mcp_operations.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", new=AsyncMock(return_value=MagicMock()), ) as create_client, patch.object( # test-quality-ok: same boundary, this is the tools/list answer the upstream would give @@ -7735,7 +7735,7 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator(): with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=alias_less_server, ), patch.object( @@ -7821,7 +7821,7 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti ), patch.object( mcp_operations.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", new=fake_create_mcp_client, ), patch.object( @@ -7888,7 +7888,7 @@ async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id(): ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=oauth_server, ), patch.object( @@ -7950,7 +7950,7 @@ async def test_execute_mcp_tool_rest_unauthorized_prefix_still_mismatches(): ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=restricted_server, ), patch.object( @@ -8011,7 +8011,7 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=None, ), patch.object( @@ -8094,7 +8094,7 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=fake_server, ), patch.object( @@ -8170,7 +8170,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), ): @@ -8225,7 +8225,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), ): await mcp_module.execute_mcp_tool( @@ -8280,7 +8280,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entr try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), ): for caller in (guarded, opted_out): @@ -8328,7 +8328,7 @@ async def test_execute_mcp_tool_runs_the_longer_colliding_operation_and_hands_ho try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), ): result = await mcp_module.execute_mcp_tool( @@ -8374,7 +8374,7 @@ async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation try: with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=petstore), patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check), ): result = await mcp_module.execute_mcp_tool( @@ -8404,7 +8404,7 @@ async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hoo upstream = AsyncMock() upstream.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False) proxy_logging = _mock_mcp_proxy_logging() - proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging.create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value={}) proxy_logging.during_call_hook = AsyncMock(return_value=None) @@ -8413,7 +8413,7 @@ async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hoo ) with ( - patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=upstream)), + patch.object(manager, "create_mcp_client", new=AsyncMock(return_value=upstream)), patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools), patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), ): @@ -8429,7 +8429,7 @@ async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hoo assert fetch_tools.await_count == 1 assert upstream.call_tool.await_count == 1 assert result.content[0].text == "ok" - hook_kwargs = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0] + hook_kwargs = proxy_logging.create_mcp_request_object_from_kwargs.call_args.args[0] assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None) assert server.server_id not in manager._listed_tools_by_server_id @@ -8460,7 +8460,7 @@ async def test_fetch_pinnable_tool_catalog_records_no_listed_catalog_for_the_adm ) with ( - patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock())), + patch.object(manager, "create_mcp_client", new=AsyncMock(return_value=MagicMock())), patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools), patch("litellm.proxy.proxy_server.proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())), ): @@ -8523,7 +8523,7 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=None, ), patch.object( @@ -8603,7 +8603,7 @@ async def test_execute_mcp_tool_rest_prefix_retry_resolution_still_enforces_serv ), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", side_effect=resolve_only_when_requested_prefix_added, ), patch.object( @@ -8671,7 +8671,7 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_unknown_name_fails_ with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups", + "MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ): @@ -8730,7 +8730,7 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_known_alias_returns with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups", + "MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ): @@ -8763,7 +8763,7 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unk with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups", + "MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ): @@ -8796,7 +8796,7 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_access_group_resolv with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups", + "MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=["id-b"], ): @@ -9630,9 +9630,9 @@ def test_redact_mcp_resource_url_strips_credentials(url, expected): """The MCP tool-call log records the upstream resource, so the URL must be redacted to scheme+host+path: userinfo, query string, and fragment (which can carry embedded tokens or secret parameters) must never reach spend-log metadata or logging callbacks.""" - from litellm.proxy._experimental.mcp_server.server import _redact_mcp_resource_url + from litellm.proxy._experimental.mcp_server.server import redact_mcp_resource_url - assert _redact_mcp_resource_url(url) == expected + assert redact_mcp_resource_url(url) == expected @pytest.mark.asyncio @@ -9721,7 +9721,7 @@ def _managed_tool_returning(server, upstream_result, proxy_logging_mock): return_value=[server.server_id], ), patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server), - patch.object(global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(global_mcp_server_manager, "get_mcp_server_from_tool_name", return_value=server), patch.object(global_mcp_server_manager, "server_owning_tool_name_prefix", return_value=server), patch( "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers_from_mcp_server_names", @@ -9898,7 +9898,7 @@ async def test_aggregate_listing_reports_per_server_outcomes(): return [tool1] raise MCPServerListError(ServerListFault(tag="upstream_error", status_code=500), server.name) - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with patch( "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", @@ -10747,7 +10747,7 @@ async def test_list_tools_injects_byok_credential_for_non_oauth2_auth_types(auth mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=[server.server_id]) mock_manager.get_mcp_server_by_id = MagicMock(return_value=server) mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) - mock_manager._get_tools_from_server = mock_get_tools_from_server + mock_manager.get_tools_from_server = mock_get_tools_from_server with ( patch( @@ -10949,7 +10949,7 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)), patch.object(operations, "function_setup", return_value=(None, None)), patch.object(proxy_server, "proxy_logging_obj", logger), - patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), + patch.object(operations.global_mcp_server_manager, "get_tools_from_server", upstream), ): with pytest.raises(HTTPException) as rejected: await operations._get_tools_from_mcp_servers( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py index 6430d5f9259..dda8f0b1044 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py @@ -615,7 +615,7 @@ class TestCredentialMergeOnUpdate: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ), patch( @@ -652,7 +652,7 @@ class TestCredentialMergeOnUpdate: ) with patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ): await update_mcp_server(mock_prisma, data, "test-user") @@ -682,7 +682,7 @@ class TestCredentialMergeOnUpdate: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ), patch( @@ -724,7 +724,7 @@ class TestCredentialMergeOnUpdate: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ), patch( @@ -767,7 +767,7 @@ class TestCredentialMergeOnUpdate: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ), patch( @@ -1001,7 +1001,7 @@ class TestRotateCredentials: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value="old-key", ), patch( @@ -1050,7 +1050,7 @@ class TestRotateCredentials: with ( patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value="old-key", ), patch( @@ -1101,7 +1101,7 @@ class TestAuthTypeSwitchClearsCredentials: ) with patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", + "litellm.proxy._experimental.mcp_server.db.get_salt_key", return_value=None, ): await update_mcp_server(mock_prisma, data, "test-user") @@ -1121,7 +1121,7 @@ class TestInheritCredentials: def test_inherits_sigv4_credentials(self): """SigV4 fields are copied from existing server to inherited credentials.""" from litellm.proxy.management_endpoints.mcp_management_endpoints import ( - _inherit_credentials_from_existing_server, + inherit_credentials_from_existing_server, ) from litellm.proxy._types import NewMCPServerRequest from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -1152,7 +1152,7 @@ class TestInheritCredentials: "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager" ) as mock_manager: mock_manager.get_mcp_server_by_id.return_value = existing - result = _inherit_credentials_from_existing_server(payload) + result = inherit_credentials_from_existing_server(payload) assert result.credentials is not None assert result.credentials["aws_access_key_id"] == "AKIAEXAMPLE" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 1c987778da6..c3a94e04b9f 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -1373,7 +1373,7 @@ async def test_mcp_tool_search_leaves_the_listed_tools_slot_empty(monkeypatch: p with ( patch.dict(manager.tool_name_to_mcp_server_name_mapping), patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), - patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "create_mcp_client", AsyncMock(return_value=object())), patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), ): try: diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py index 95be2b8b12b..097b0924b8d 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py @@ -50,7 +50,7 @@ class TestApplyToolsetScope: @pytest.mark.asyncio async def test_restricts_to_toolset_servers_and_tools(self): - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope toolset_perms = { "server-a": ["tool1", "tool2"], @@ -66,7 +66,7 @@ class TestApplyToolsetScope: mcp_servers=["server-a", "server-b", "server-c"], mcp_toolsets=["toolset-123"], ) - result = await _apply_toolset_scope(auth, "toolset-123") + result = await apply_toolset_scope(auth, "toolset-123") op = result.object_permission assert op is not None @@ -92,7 +92,7 @@ class TestApplyToolsetScope: @pytest.mark.asyncio async def test_admin_creates_object_permission_when_none(self): """Admin key with object_permission=None can access any toolset.""" - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope toolset_perms = {"server-a": ["tool1"]} with patch( @@ -105,7 +105,7 @@ class TestApplyToolsetScope: user_role=LitellmUserRoles.PROXY_ADMIN, object_permission=None, ) - result = await _apply_toolset_scope(auth, "toolset-123") + result = await apply_toolset_scope(auth, "toolset-123") op = result.object_permission assert op is not None @@ -116,7 +116,7 @@ class TestApplyToolsetScope: async def test_team_granted_toolset_is_served_to_a_key_without_its_own_grant(self): """A team key whose own row carries no toolset grant is admitted to the toolset its team holds (LIT-6029), scoped to that toolset's servers and tools.""" - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope toolset_perms = {"server-a": ["tool1"]} auth = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=None) @@ -125,7 +125,7 @@ class TestApplyToolsetScope: "global_mcp_server_manager.resolve_toolset_tool_permissions", new=AsyncMock(return_value=toolset_perms), ): - result = await _apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-123")) + result = await apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-123")) assert result.mcp_toolset_id == "toolset-123" assert result.object_permission is not None @@ -138,7 +138,7 @@ class TestApplyToolsetScope: team-granted toolset is not capped by the user's own row: the row stays intact and the toolset rides along as mcp_toolset_id (LIT-6029).""" from litellm.constants import UI_SESSION_TOKEN_TEAM_ID - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") own_row = LiteLLM_ObjectPermissionTable(object_permission_id="user-op", mcp_servers=["server-own"]) @@ -151,7 +151,7 @@ class TestApplyToolsetScope: "global_mcp_server_manager.resolve_toolset_tool_permissions", new=resolve, ): - result = await _apply_toolset_scope( + result = await apply_toolset_scope( session, "toolset-123", acting_user=AsyncMock(return_value=admitted), granted=granted ) @@ -163,13 +163,13 @@ class TestApplyToolsetScope: @pytest.mark.asyncio async def test_a_gateway_admitted_user_without_the_toolset_in_any_source_is_denied(self): - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) admitted.mcp_admitted_user_subject = True granted = AsyncMock(return_value=frozenset({"toolset-other"})) with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + await apply_toolset_scope(admitted, "toolset-123", granted=granted) assert exc_info.value.status_code == 403 granted.assert_awaited_once_with(admitted) @@ -178,7 +178,7 @@ class TestApplyToolsetScope: async def test_a_resource_scoped_admitted_user_is_denied_a_team_toolset_on_another_server(self): """A gateway bearer scoped to server-own (RFC 8707 resource) cannot open a team toolset whose servers lie outside that resource, even though the team grants it (Devin Review 4150024267).""" - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) admitted.mcp_admitted_user_subject = True @@ -192,14 +192,14 @@ class TestApplyToolsetScope: new=resolve, ): with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + await apply_toolset_scope(admitted, "toolset-123", granted=granted) assert exc_info.value.status_code == 403 resolve.assert_awaited_once_with(toolset_ids=["toolset-123"], requires_fresh_policy=True) @pytest.mark.asyncio async def test_a_resource_scoped_admitted_user_opens_a_toolset_inside_its_resource(self): - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) admitted.mcp_admitted_user_subject = True @@ -211,7 +211,7 @@ class TestApplyToolsetScope: "global_mcp_server_manager.resolve_toolset_tool_permissions", new=resolve, ): - result = await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + result = await apply_toolset_scope(admitted, "toolset-123", granted=granted) assert result.mcp_toolset_id == "toolset-123" assert result.mcp_session_resource_server_id == "server-team" @@ -219,12 +219,12 @@ class TestApplyToolsetScope: @pytest.mark.asyncio async def test_team_grant_for_another_toolset_does_not_admit_a_key_to_this_one(self): - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope auth = _make_auth(mcp_toolsets=[]) auth.team_id = "team-a" with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-other")) + await apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-other")) assert exc_info.value.status_code == 403 @@ -233,11 +233,11 @@ class TestApplyToolsetScope: """Non-admin key with object_permission=None is denied (no grants configured).""" from starlette.exceptions import HTTPException - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope auth = UserAPIKeyAuth(api_key="sk-test", object_permission=None) with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(auth, "toolset-123") + await apply_toolset_scope(auth, "toolset-123") assert exc_info.value.status_code == 403 @pytest.mark.asyncio @@ -248,7 +248,7 @@ class TestApplyToolsetScope: toolset path, which replaces mcp_servers and would drop the sentinel.""" from starlette.exceptions import HTTPException - from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + from litellm.proxy._experimental.mcp_server.server import apply_toolset_scope op = LiteLLM_ObjectPermissionTable( object_permission_id="test", @@ -266,7 +266,7 @@ class TestApplyToolsetScope: new=resolve, ): with pytest.raises(HTTPException) as exc_info: - await _apply_toolset_scope(auth, "toolset-123") + await apply_toolset_scope(auth, "toolset-123") assert exc_info.value.status_code == 403 resolve.assert_not_awaited() @@ -786,17 +786,17 @@ class TestMCPActiveToolsetContextVar: """Tests for _mcp_active_toolset_id ContextVar — clients cannot inject it.""" def test_contextvar_default_is_none(self): - from litellm.proxy._experimental.mcp_server.server import _mcp_active_toolset_id + from litellm.proxy._experimental.mcp_server.server import mcp_active_toolset_id - assert _mcp_active_toolset_id.get() is None + assert mcp_active_toolset_id.get() is None def test_contextvar_set_and_reset(self): - from litellm.proxy._experimental.mcp_server.server import _mcp_active_toolset_id + from litellm.proxy._experimental.mcp_server.server import mcp_active_toolset_id - token = _mcp_active_toolset_id.set("toolset-abc") - assert _mcp_active_toolset_id.get() == "toolset-abc" - _mcp_active_toolset_id.reset(token) - assert _mcp_active_toolset_id.get() is None + token = mcp_active_toolset_id.set("toolset-abc") + assert mcp_active_toolset_id.get() == "toolset-abc" + mcp_active_toolset_id.reset(token) + assert mcp_active_toolset_id.get() is None @pytest.mark.asyncio async def test_client_header_is_stripped_in_scope(self): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_oauth2_token_cache.py b/tests/unit/proxy/_experimental/mcp_server/test_oauth2_token_cache.py index 30d0f17a099..6eb93b090e9 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_oauth2_token_cache.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_oauth2_token_cache.py @@ -237,44 +237,44 @@ def test_storage_ttl_capped_at_token_lifetime(): while the stored refresh_token sat unused because refresh only runs on the DB read-through.""" from litellm.constants import MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None, token_storage_ttl_seconds=604800) - assert _compute_per_user_token_ttl(server, expires_in=86400) == 86400 - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS + assert compute_per_user_token_ttl(server, expires_in=86400) == 86400 - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS def test_storage_ttl_shorter_than_token_lifetime_wins(): """A configured TTL below the token lifetime is the operative value: the knob's purpose is to force earlier DB re-checks (staleness backstop), so the shorter side must win the min().""" from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None, token_storage_ttl_seconds=3600) - assert _compute_per_user_token_ttl(server, expires_in=86400) == 3600 + assert compute_per_user_token_ttl(server, expires_in=86400) == 3600 def test_storage_ttl_verbatim_when_token_lifetime_unknown(): """With no expires_in from the upstream there is nothing to cap against, so the configured TTL applies as-is (matching the pre-cap behavior for lifetime-less tokens).""" from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None, token_storage_ttl_seconds=604800) - assert _compute_per_user_token_ttl(server, expires_in=None) == 604800 + assert compute_per_user_token_ttl(server, expires_in=None) == 604800 def test_storage_ttl_floors_at_one_second_for_nearly_dead_token(): """A token already inside the expiry buffer yields the 1-second floor, not zero or a negative TTL, mirroring the floor the default (unconfigured) path has always had.""" from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None, token_storage_ttl_seconds=3600) - assert _compute_per_user_token_ttl(server, expires_in=30) == 1 + assert compute_per_user_token_ttl(server, expires_in=30) == 1 def test_default_ttl_paths_unchanged_without_storage_ttl(): @@ -285,12 +285,12 @@ def test_default_ttl_paths_unchanged_without_storage_ttl(): MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS, ) from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( - _compute_per_user_token_ttl, + compute_per_user_token_ttl, ) server = _server(oauth2_flow=None) - assert _compute_per_user_token_ttl(server, expires_in=86400) == 86400 - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS - assert _compute_per_user_token_ttl(server, expires_in=None) == MCP_PER_USER_TOKEN_DEFAULT_TTL + assert compute_per_user_token_ttl(server, expires_in=86400) == 86400 - MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS + assert compute_per_user_token_ttl(server, expires_in=None) == MCP_PER_USER_TOKEN_DEFAULT_TTL @pytest.mark.asyncio diff --git a/tests/unit/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py b/tests/unit/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py index 6b0211c3866..9f3db632db1 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -20,9 +20,9 @@ from respx import MockRouter from litellm.types.mcp import MCPAuth, MCPAuthType from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, - _request_resolved_auth_headers, + request_auth_header, + request_extra_headers, + request_resolved_auth_headers, _request_upstream_url, _resolve_param_list, _resolve_ref, @@ -99,7 +99,7 @@ async def test_authorization_validates_credentials_before_http( "/echo", "get", {}, "https://upstream.example", auth_type=auth_type, ) destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") - caller_token: Final = _request_auth_header.set(value) + caller_token: Final = request_auth_header.set(value) try: if accepted: assert await tool() == TextResult("authenticated") @@ -111,7 +111,7 @@ async def test_authorization_validates_credentials_before_http( assert exc.value.status_code == 500 assert destination.call_count == 0 finally: - _request_auth_header.reset(caller_token) + request_auth_header.reset(caller_token) @pytest.mark.asyncio @@ -132,9 +132,9 @@ async def test_static_auth_validates_headers_after_existing_precedence( ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") - caller_token: Final = _request_auth_header.set(caller) - extra_token: Final = _request_extra_headers.set(forwarded) - resolved_token: Final = _request_resolved_auth_headers.set(resolved) + caller_token: Final = request_auth_header.set(caller) + extra_token: Final = request_extra_headers.set(forwarded) + resolved_token: Final = request_resolved_auth_headers.set(resolved) try: if expected is None: with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc: @@ -146,9 +146,9 @@ async def test_static_auth_validates_headers_after_existing_precedence( assert destination.call_count == 1 assert destination.calls.last.request.headers["authorization"] == expected finally: - _request_auth_header.reset(caller_token) - _request_extra_headers.reset(extra_token) - _request_resolved_auth_headers.reset(resolved_token) + request_auth_header.reset(caller_token) + request_extra_headers.reset(extra_token) + request_resolved_auth_headers.reset(resolved_token) @pytest.mark.asyncio @@ -204,13 +204,13 @@ async def test_static_validation_preserves_no_auth_and_resolved_oauth( monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") tool: Final = create_tool_function("/echo", "get", {}, "https://upstream.example", auth_type=auth_type) destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="echo") - token: Final = _request_resolved_auth_headers.set(resolved) + token: Final = request_resolved_auth_headers.set(resolved) try: assert await tool() == TextResult("echo") assert destination.call_count == 1 assert destination.calls.last.request.headers.get("authorization") == (resolved or {}).get("Authorization") finally: - _request_resolved_auth_headers.reset(token) + request_resolved_auth_headers.reset(token) def _create_mock_client(method: str, response_text: str, status_code: int = 200) -> AsyncMock: @@ -1214,11 +1214,11 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_extra_headers.set({"X-TOKEN": "secret-value"}) + token = request_extra_headers.set({"X-TOKEN": "secret-value"}) try: result = await func() finally: - _request_extra_headers.reset(token) + request_extra_headers.reset(token) assert result == TextResult("ok") call_args = async_client.get.call_args @@ -1265,11 +1265,11 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("post", "created") mock_client.return_value = async_client - token = _request_extra_headers.set({"X-TOKEN": "dynamic-value"}) + token = request_extra_headers.set({"X-TOKEN": "dynamic-value"}) try: result = await func() finally: - _request_extra_headers.reset(token) + request_extra_headers.reset(token) assert result == TextResult("created") call_args = async_client.post.call_args @@ -1293,11 +1293,11 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_extra_headers.set({"X-Tenant": "caller-spoofed"}) + token = request_extra_headers.set({"X-Tenant": "caller-spoofed"}) try: result = await func() finally: - _request_extra_headers.reset(token) + request_extra_headers.reset(token) assert result == TextResult("ok") call_args = async_client.get.call_args @@ -1321,11 +1321,11 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_extra_headers.set({"x-tenant": "caller-spoofed"}) + token = request_extra_headers.set({"x-tenant": "caller-spoofed"}) try: result = await func() finally: - _request_extra_headers.reset(token) + request_extra_headers.reset(token) assert result == TextResult("ok") call_args = async_client.get.call_args @@ -1349,15 +1349,15 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "secure-data") mock_client.return_value = async_client - extra_token = _request_extra_headers.set( + extra_token = request_extra_headers.set( {"Authorization": "Bearer extra", "X-TOKEN": "token-value"} ) - auth_token = _request_auth_header.set("Bearer byok-credential") + auth_token = request_auth_header.set("Bearer byok-credential") try: result = await func() finally: - _request_auth_header.reset(auth_token) - _request_extra_headers.reset(extra_token) + request_auth_header.reset(auth_token) + request_extra_headers.reset(extra_token) assert result == TextResult("secure-data") call_args = async_client.get.call_args @@ -1380,8 +1380,8 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_extra_headers.set({"X-TOKEN": "first-call"}) - _request_extra_headers.reset(token) + token = request_extra_headers.set({"X-TOKEN": "first-call"}) + request_extra_headers.reset(token) await func() @@ -1409,15 +1409,15 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "secure-data") mock_client.return_value = async_client - extra_token = _request_extra_headers.set({"Authorization": "Bearer caller-forwarded"}) - auth_token = _request_auth_header.set("Bearer byok-credential") - resolved_token = _request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) + extra_token = request_extra_headers.set({"Authorization": "Bearer caller-forwarded"}) + auth_token = request_auth_header.set("Bearer byok-credential") + resolved_token = request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) try: result = await func() finally: - _request_auth_header.reset(auth_token) - _request_extra_headers.reset(extra_token) - _request_resolved_auth_headers.reset(resolved_token) + request_auth_header.reset(auth_token) + request_extra_headers.reset(extra_token) + request_resolved_auth_headers.reset(resolved_token) assert result == TextResult("secure-data") headers_sent = async_client.get.call_args[1]["headers"] @@ -1439,8 +1439,8 @@ class TestRequestExtraHeaders: async_client = _create_mock_client("get", "ok") mock_client.return_value = async_client - token = _request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) - _request_resolved_auth_headers.reset(token) + token = request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) + request_resolved_auth_headers.reset(token) await func() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index c570d498f44..ffa25d63391 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -56,7 +56,7 @@ async def test_openapi_local_tool_runs_pre_call_tool_check(): with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=fake_server, ), patch.object( @@ -145,7 +145,7 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises(): with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=fake_server, ), patch.object( @@ -207,10 +207,10 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable(): # `_get_mcp_server_from_tool_name` returns None — no server context. with ( - patch.object(mcp_operations, "_resolve_openapi_tool_auth", new=resolve_auth), + patch.object(mcp_operations, "resolve_openapi_tool_auth", new=resolve_auth), patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=None, ), patch.object( @@ -258,7 +258,7 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): reads. Kills the mutant that deletes the resolve_openapi_upstream_auth call in server.py.""" from litellm.proxy._experimental.mcp_server import server as mcp_module from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( StaticHeaderAuth, @@ -290,13 +290,13 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): captured: dict = {} async def handle_local(_name, _arguments, _wire_compat): - captured["resolved"] = _request_resolved_auth_headers.get() + captured["resolved"] = request_resolved_auth_headers.get() return CallToolResult(content=[], is_error=False) with ( patch.object( mcp_operations.global_mcp_server_manager, - "_get_mcp_server_from_tool_name", + "get_mcp_server_from_tool_name", return_value=oauth_server, ), patch.object( @@ -332,7 +332,7 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): ) assert captured["resolved"] == {"Authorization": "Bearer stored-user-token"} - assert _request_resolved_auth_headers.get() is None + assert request_resolved_auth_headers.get() is None @@ -605,7 +605,7 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc """ from litellm.proxy._experimental.mcp_server import server as mcp_module from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, + request_auth_header, ) server = _spec_path_server() @@ -618,11 +618,11 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc return None, kwargs["forwarded_headers"] async def capture_local(_name, _arguments, _wire_compat): - captured["injected"] = _request_auth_header.get() + captured["injected"] = request_auth_header.get() return CallToolResult(content=[], is_error=False) async def capture_openapi_handler(_server, _name, _arguments, _wire_compat): - captured["injected"] = _request_auth_header.get() + captured["injected"] = request_auth_header.get() return CallToolResult(content=[], is_error=False) manager = mcp_operations.global_mcp_server_manager @@ -637,7 +637,7 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc fake_tool.input_schema = {"type": "object"} fake_tool.server_id = server.server_id with ( - patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(manager, "get_mcp_server_from_tool_name", return_value=server), patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool), patch( "litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool", @@ -671,7 +671,7 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc assert captured["resolver_credential"] == {"Authorization": OPENAPI_PER_SERVER_TOKEN} assert captured["injected"] == OPENAPI_PER_SERVER_TOKEN - assert _request_auth_header.get() is None + assert request_auth_header.get() is None @pytest.mark.parametrize("failure", ["auth", "other"]) @@ -723,7 +723,7 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st user = UserAPIKeyAuth(api_key="sk-user", user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER.value) with ( - patch.object(mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(mcp_operations.global_mcp_server_manager, "get_mcp_server_from_tool_name", return_value=server), patch.object(mcp_operations.global_mcp_server_manager, "pre_call_tool_check", new=AsyncMock(return_value={})), patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool), patch.object( @@ -798,16 +798,16 @@ def test_the_openapi_arm_installs_the_guard_when_a_credential_rides_a_custom_slo hook alone passes even if this arm never installs it. """ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, _upstream_client, ) - token = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) + token = request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) try: client = _upstream_client() assert client.client.event_hooks["request"], "custom slot must install a redirect guard" finally: - _request_resolved_auth_headers.reset(token) + request_resolved_auth_headers.reset(token) def test_the_guarded_client_is_reused_rather_than_built_per_call(): @@ -816,15 +816,15 @@ def test_the_guarded_client_is_reused_rather_than_built_per_call(): have to come from the shared cache. """ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, _upstream_client, ) - token = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) + token = request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) try: assert _upstream_client() is _upstream_client() finally: - _request_resolved_auth_headers.reset(token) + request_resolved_auth_headers.reset(token) @pytest.mark.asyncio @@ -836,11 +836,11 @@ async def test_the_shared_guard_reads_the_url_from_the_request_context(): from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _drop_credential_across_origin, - _request_resolved_auth_headers, + request_resolved_auth_headers, _request_upstream_url, ) - creds = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) + creds = request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"}) url = _request_upstream_url.set("https://api.example.com/v1/things") try: same = httpx.Request("POST", "https://api.example.com/v1/other", headers={"esb-oauth": "Bearer m"}) @@ -852,7 +852,7 @@ async def test_the_shared_guard_reads_the_url_from_the_request_context(): assert "esb-oauth" not in foreign.headers finally: _request_upstream_url.reset(url) - _request_resolved_auth_headers.reset(creds) + request_resolved_auth_headers.reset(creds) @pytest.mark.parametrize("resolved", [{"Authorization": "Bearer minted"}, {}, None]) @@ -860,16 +860,16 @@ def test_the_openapi_arm_keeps_the_shared_client_when_no_guard_is_needed(resolve # Authorization is already stripped across origins by the HTTP client, so taking the guarded # path for it would give up the shared connection pool for nothing. from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_resolved_auth_headers, + request_resolved_auth_headers, _upstream_client, ) - token = _request_resolved_auth_headers.set(resolved) + token = request_resolved_auth_headers.set(resolved) try: client = _upstream_client() assert not client.client.event_hooks.get("request") finally: - _request_resolved_auth_headers.reset(token) + request_resolved_auth_headers.reset(token) @pytest.mark.asyncio diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index 089426eb5cb..d9997ec0fce 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -22,7 +22,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token from litellm.types.mcp import MCPAuth, MCPTransport @@ -294,7 +294,7 @@ def _catalog_case(method): def _mcp_rate_limited_proxy_logging() -> ProxyLogging: proxy_logging: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache()) - proxy_logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + proxy_logging.proxy_hook_mapping["parallel_request_limiter"] = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(DualCache()) ) return proxy_logging @@ -324,7 +324,7 @@ async def test_mcp_server_rpm_limits_every_catalog_operation(operation: str) -> ) caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-catalog-rpm")) operation_to_manager_method: Final = { - "tools/list": "_get_tools_from_server", + "tools/list": "get_tools_from_server", "prompts/list": "get_prompts_from_server", "resources/list": "get_resources_from_server", "resources/templates/list": "get_resource_templates_from_server", @@ -433,7 +433,7 @@ async def test_tools_call_warmup_does_not_consume_mcp_server_rpm() -> None: patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), patch.object(operations.global_mcp_server_manager, "server_exposes_tool", return_value=False), patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), - patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), + patch.object(operations.global_mcp_server_manager, "get_tools_from_server", upstream), ): await operations._list_tools_before_first_call( server=server, @@ -502,8 +502,8 @@ async def test_tools_call_pre_call_hook_rejection_does_not_enforce_mcp_server_rp ) rate_limit_error: Final = ProxyRateLimitError(detail="ordinary key rate limit") proxy_logging: Final = MagicMock() - proxy_logging._create_mcp_request_object_from_kwargs.return_value = {} - proxy_logging._convert_mcp_to_llm_format.return_value = {} + proxy_logging.create_mcp_request_object_from_kwargs.return_value = {} + proxy_logging.convert_mcp_to_llm_format.return_value = {} proxy_logging.pre_call_hook = AsyncMock(side_effect=rate_limit_error) proxy_logging.enforce_mcp_server_rate_limits = AsyncMock() @@ -714,7 +714,7 @@ async def test_tool_listing_returns_empty_result_without_dispatch_for_unavailabl upstream = AsyncMock() with ( patch.object(operations, "_get_allowed_mcp_servers", allowed), - patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), + patch.object(operations.global_mcp_server_manager, "get_tools_from_server", upstream), ): result = await GatewayOperations().execute(ListToolsRequest(), prepare_context()) assert result.tools == [] @@ -1142,7 +1142,7 @@ async def test_list_mcp_tools_records_the_catalog_only_when_asked( upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})] with ( patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), - patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "create_mcp_client", AsyncMock(return_value=object())), patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), patch.dict(manager.tool_name_to_mcp_server_name_mapping), ): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py b/tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py index ac453df8fa5..efddd1beef3 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_per_user_oauth_cache.py @@ -21,7 +21,7 @@ for _mod in ("orjson",): from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( # noqa: E402 MCPPerUserTokenCache, - _compute_per_user_token_ttl, + compute_per_user_token_ttl, mcp_per_user_token_cache, ) from litellm.types.mcp import MCPAuth, MCPTransport # noqa: E402 @@ -215,25 +215,25 @@ class TestValidateTokenResponse: class TestComputePerUserTokenTtl: def test_uses_server_override_when_set(self): server = _make_server(token_storage_ttl_seconds=7200) - assert _compute_per_user_token_ttl(server, expires_in=99999) == 7200 + assert compute_per_user_token_ttl(server, expires_in=99999) == 7200 def test_uses_expires_in_minus_buffer(self): server = _make_server() # Default buffer is 60s - ttl = _compute_per_user_token_ttl(server, expires_in=3600) + ttl = compute_per_user_token_ttl(server, expires_in=3600) assert ttl == 3600 - 60 def test_minimum_ttl_is_1(self): server = _make_server() # expires_in smaller than buffer → clamp to 1 - ttl = _compute_per_user_token_ttl(server, expires_in=30) + ttl = compute_per_user_token_ttl(server, expires_in=30) assert ttl == 1 def test_default_ttl_when_expires_in_none(self): from litellm.constants import MCP_PER_USER_TOKEN_DEFAULT_TTL server = _make_server() - ttl = _compute_per_user_token_ttl(server, expires_in=None) + ttl = compute_per_user_token_ttl(server, expires_in=None) assert ttl == MCP_PER_USER_TOKEN_DEFAULT_TTL diff --git a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py index c748bdc6a6b..791e6d18ea8 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -218,7 +218,7 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, ) @@ -243,7 +243,7 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, ) @@ -271,7 +271,7 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", hanging_create_client, ) @@ -350,13 +350,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -400,13 +400,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -451,13 +451,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -492,13 +492,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -546,13 +546,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -600,13 +600,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -648,13 +648,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -692,13 +692,13 @@ class TestExecuteWithMcpClient: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_build_stdio_env", + "build_stdio_env", fake_build_stdio_env, raising=False, ) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_create_mcp_client", + "create_mcp_client", fake_create_client, raising=False, ) @@ -900,7 +900,7 @@ class TestTestToolsList: monkeypatch.setattr( auth_mcp.MCPRequestHandler, - "_get_oauth2_headers_from_headers", + "get_oauth2_headers_from_headers", staticmethod(fake_oauth), raising=False, ) @@ -1082,7 +1082,7 @@ class TestTestToolsList: monkeypatch.setattr( auth_mcp.MCPRequestHandler, - "_get_oauth2_headers_from_headers", + "get_oauth2_headers_from_headers", staticmethod(fake_oauth), raising=False, ) @@ -1136,7 +1136,7 @@ class TestTestToolsList: monkeypatch.setattr( auth_mcp.MCPRequestHandler, - "_get_oauth2_headers_from_headers", + "get_oauth2_headers_from_headers", staticmethod(lambda headers: oauth_headers), raising=False, ) @@ -1233,7 +1233,7 @@ class TestListToolsRestAPI: monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging) monkeypatch.setattr(manager, "get_allowed_mcp_servers", allowed_servers) monkeypatch.setattr(manager, "filter_server_ids_by_ip_with_info", lambda ids, _ip: (ids, 0)) - monkeypatch.setattr(manager, "_get_tools_from_server", upstream) + monkeypatch.setattr(manager, "get_tools_from_server", upstream) with pytest.raises(HTTPException) as error: await rest_endpoints.list_tool_rest_api( @@ -1274,7 +1274,7 @@ class TestListToolsRestAPI: monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging) monkeypatch.setattr(manager, "get_allowed_mcp_servers", allowed_servers) monkeypatch.setattr(manager, "filter_server_ids_by_ip_with_info", lambda ids, _ip: (ids, 0)) - monkeypatch.setattr(manager, "_get_tools_from_server", upstream) + monkeypatch.setattr(manager, "get_tools_from_server", upstream) result: Final = await rest_endpoints.list_tool_rest_api( _build_request(path="/mcp-rest/tools/list", method="GET"), @@ -1532,7 +1532,7 @@ class TestListToolsRestAPI: fake_get_toolset_by_name_cached, raising=False, ) - monkeypatch.setattr(rest_endpoints, "_apply_toolset_scope", fake_apply_toolset_scope, raising=False) + monkeypatch.setattr(rest_endpoints, "apply_toolset_scope", fake_apply_toolset_scope, raising=False) monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", @@ -2319,7 +2319,7 @@ class TestListToolsRestAPI: ) monkeypatch.setattr( rest_endpoints, - "_apply_toolset_scope", + "apply_toolset_scope", fake_apply_toolset_scope, raising=False, ) @@ -3497,7 +3497,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3550,7 +3550,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3592,7 +3592,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3639,7 +3639,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3688,7 +3688,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -3739,7 +3739,7 @@ class TestGetToolsForSingleServer: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", fake_get_tools_from_server, raising=False, ) @@ -4518,7 +4518,7 @@ class TestRestListToolsetFiltering: monkeypatch.setattr( rest_endpoints.global_mcp_server_manager, - "_get_tools_from_server", + "get_tools_from_server", AsyncMock(return_value=upstream_tools), ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py b/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py index 51b97f7c8c5..4b9214916a8 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py @@ -43,11 +43,11 @@ class FakeMCPServerManager: self.name_lookup_spy(server_name, client_ip) return self.servers_by_name.get(server_name) - def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: + def is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: self.ip_filter_spy(server, client_ip) return self.ip_accessible - def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: + def build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: return LiteLLM_MCPServerTable( server_id=server.server_id, alias=server.alias, @@ -132,7 +132,7 @@ async def test_temp_resolution_precedes_db_and_registry() -> None: ) assert resolved == ResolvedMCPServer( - table=manager._build_mcp_server_table(temporary_server), + table=manager.build_mcp_server_table(temporary_server), runtime=temporary_server, source="temp", ) @@ -179,7 +179,7 @@ async def test_registry_id_resolution_precedes_name() -> None: ) assert resolved == ResolvedMCPServer( - table=manager._build_mcp_server_table(server), + table=manager.build_mcp_server_table(server), runtime=server, source="registry", ) @@ -244,7 +244,7 @@ async def test_db_lookup_none_skips_db_and_returns_registry_source() -> None: resolved: Final = await resolve_mcp_server(server.server_id, manager=manager, db_lookup=None) assert resolved == ResolvedMCPServer( - table=manager._build_mcp_server_table(server), + table=manager.build_mcp_server_table(server), runtime=server, source="registry", ) @@ -286,7 +286,7 @@ async def test_non_admin_temp_resolution_is_denied_before_allowed_lookup() -> No server: Final = _runtime_server() manager: Final = _manager(allowed_server_ids=(server.server_id,)) resolved: Final = ResolvedMCPServer( - table=manager._build_mcp_server_table(server), + table=manager.build_mcp_server_table(server), runtime=server, source="temp", ) @@ -384,7 +384,7 @@ async def test_catalog_visibility_never_opens_temporary_setup_to_non_admins( ) -> None: server: Final = _runtime_server() manager: Final = _manager() - resolved: Final = ResolvedMCPServer(manager._build_mcp_server_table(server), server, source) + resolved: Final = ResolvedMCPServer(manager.build_mcp_server_table(server), server, source) operation: Final = authorize_mcp_server( resolved, _auth(), diff --git a/tests/unit/proxy/agent_endpoints/auth/test_agent_permission_handler.py b/tests/unit/proxy/agent_endpoints/auth/test_agent_permission_handler.py index a1a022fdd35..b5fa1cbc53d 100644 --- a/tests/unit/proxy/agent_endpoints/auth/test_agent_permission_handler.py +++ b/tests/unit/proxy/agent_endpoints/auth/test_agent_permission_handler.py @@ -526,7 +526,7 @@ class TestAgentRequestHandler: AgentRequestHandler, "_get_key_object_permission", return_value=None ): with patch( - "litellm.proxy.auth.auth_checks._get_agent_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_agent_ids_from_access_groups", new_callable=AsyncMock, return_value=["agent-from-ag-1", "agent-from-ag-2"], ): @@ -558,7 +558,7 @@ class TestAgentRequestHandler: mock_user_auth.object_permission = mock_permission with patch( - "litellm.proxy.auth.auth_checks._get_agent_ids_from_access_groups", + "litellm.proxy.auth.auth_checks.get_agent_ids_from_access_groups", new_callable=AsyncMock, return_value=["agent-from-ag"], ): @@ -966,7 +966,7 @@ async def test_managed_target_rechecks_authoritative_key_after_peer_revocation( warm.object_permission = None warm.access_group_ids = ["old-group"] from litellm.proxy.auth import auth_checks - monkeypatch.setattr(auth_checks, "_get_agent_ids_from_access_groups", AsyncMock(return_value=["target"])) + monkeypatch.setattr(auth_checks, "get_agent_ids_from_access_groups", AsyncMock(return_value=["target"])) if change == "deleted": client.get_data.return_value = None if change == "outage": diff --git a/tests/unit/proxy/anthropic_endpoints/test_endpoints.py b/tests/unit/proxy/anthropic_endpoints/test_endpoints.py index 801f61aa498..c0e7a4fdb47 100644 --- a/tests/unit/proxy/anthropic_endpoints/test_endpoints.py +++ b/tests/unit/proxy/anthropic_endpoints/test_endpoints.py @@ -107,7 +107,7 @@ class TestBlockedResponseUsage: ) with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})), + patch.object(ep, "read_request_body", new=AsyncMock(return_value={})), patch.object( ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", @@ -147,7 +147,7 @@ class TestProxyExceptionAnthropicEnvelope: request.headers = {"x-request-id": "req_test_6468"} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})), + patch.object(ep, "read_request_body", new=AsyncMock(return_value={})), patch.object( ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", @@ -278,7 +278,7 @@ class TestHttpExceptionDictDetail: request.headers = {} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: endpoint reads the body via a module function; no injection seam + patch.object(ep, "read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: endpoint reads the body via a module function; no injection seam patch.object( # test-quality-ok: the guardrail raise happens deep inside this call; the test targets the endpoint's except block ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", @@ -324,7 +324,7 @@ class TestFailureHookRequestData: request.headers = {} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), + patch.object(ep, "read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process), patch.object(proxy_server, "proxy_logging_obj") as mock_logging, ): @@ -375,7 +375,7 @@ class TestErrorLogCarriesCallId: request.headers = {} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam + patch.object(ep, "read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process), # test-quality-ok: the provider failure happens inside this call; the test targets the endpoint's except block patch.object(proxy_server, "proxy_logging_obj") as mock_logging, # test-quality-ok: module global imported at call time; no injection seam caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"), @@ -408,7 +408,7 @@ class TestErrorLogCarriesCallId: request.headers = {} with ( - patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam + patch.object(ep, "read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), # test-quality-ok: endpoint reads the body via a module function; no injection seam patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process), # test-quality-ok: the proxy shaped failure happens inside this call; the test targets the endpoint's except block patch.object(proxy_server, "proxy_logging_obj") as mock_logging, # test-quality-ok: module global imported at call time; no injection seam ): @@ -437,7 +437,7 @@ class TestErrorLogCarriesCallId: with ( patch.object( # test-quality-ok: endpoint reads the body via a module function; no injection seam ep, - "_read_request_body", + "read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet", "messages": [{"role": "user", "content": "hi"}]}), ), patch.object(proxy_server, "token_counter", new=AsyncMock(side_effect=RuntimeError("tokenizer down"))), # test-quality-ok: module global imported at call time; the test targets the endpoint's except block diff --git a/tests/unit/proxy/auth/test_auth_checks.py b/tests/unit/proxy/auth/test_auth_checks.py index b698a3d02fa..dfd00a7458f 100644 --- a/tests/unit/proxy/auth/test_auth_checks.py +++ b/tests/unit/proxy/auth/test_auth_checks.py @@ -32,8 +32,8 @@ from litellm.proxy._types import ( from litellm.proxy.utils import PrismaClient from litellm.proxy.auth.auth_checks import ( can_team_access_model, - _is_model_cost_zero, - _virtual_key_soft_budget_check, + is_model_cost_zero, + virtual_key_soft_budget_check, _team_soft_budget_check, ) from litellm.proxy.utils import ProxyLogging @@ -91,7 +91,7 @@ async def test_check_end_user_budget(customer_spend, customer_budget): Note: Budget enforcement for end users happens in common_checks() via _check_end_user_budget(), not in get_end_user_object(). """ - from litellm.proxy.auth.auth_checks import _check_end_user_budget + from litellm.proxy.auth.auth_checks import check_end_user_budget _budget = LiteLLM_BudgetTable(max_budget=customer_budget) end_user_obj = LiteLLM_EndUserTable( @@ -104,14 +104,14 @@ async def test_check_end_user_budget(customer_spend, customer_budget): should_exceed = customer_spend > customer_budget if not should_exceed: - await _check_end_user_budget( + await check_end_user_budget( end_user_obj=end_user_obj, route="/v1/chat/completions", ) return with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _check_end_user_budget( + await check_end_user_budget( end_user_obj=end_user_obj, route="/v1/chat/completions", ) @@ -476,7 +476,7 @@ async def test_virtual_key_max_budget_check( 1. Triggers budget alert for all cases 2. Raises BudgetExceededError when spend >= max_budget """ - from litellm.proxy.auth.auth_checks import _virtual_key_max_budget_check + from litellm.proxy.auth.auth_checks import virtual_key_max_budget_check # Setup test data valid_token = UserAPIKeyAuth( @@ -508,7 +508,7 @@ async def test_virtual_key_max_budget_check( if expect_budget_error: with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -516,7 +516,7 @@ async def test_virtual_key_max_budget_check( assert exc_info.value.current_cost == token_spend assert exc_info.value.max_budget == max_budget else: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -633,7 +633,7 @@ async def test_virtual_key_soft_budget_check(spend, soft_budget, expect_alert): proxy_logging_obj = MockProxyLogging() - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -970,7 +970,7 @@ async def test_can_key_call_model_with_aliases(model, alias_map, expect_to_work) @pytest.mark.asyncio async def test_cache_access_object(): """Test _cache_access_object stores access group in cache with correct key.""" - from litellm.proxy.auth.auth_checks import _cache_access_object + from litellm.proxy.auth.auth_checks import cache_access_object from litellm.proxy._types import LiteLLM_AccessGroupTable cache = DualCache() @@ -980,7 +980,7 @@ async def test_cache_access_object(): access_group_name="test-group", access_model_names=["gpt-4"], ) - await _cache_access_object( + await cache_access_object( access_group_id=ag_id, access_group_table=ag_table, user_api_key_cache=cache, @@ -998,7 +998,7 @@ async def test_cache_access_object(): @pytest.mark.asyncio async def test_delete_cache_access_object(): """Test _delete_cache_access_object removes access group from in-memory cache.""" - from litellm.proxy.auth.auth_checks import _delete_cache_access_object + from litellm.proxy.auth.auth_checks import delete_cache_access_object from litellm.proxy._types import LiteLLM_AccessGroupTable cache = DualCache() @@ -1008,7 +1008,7 @@ async def test_delete_cache_access_object(): access_group_name="to-delete", ) await cache.async_set_cache(key=f"access_group_id:{ag_id}", value=ag_table, ttl=60) - await _delete_cache_access_object(access_group_id=ag_id, user_api_key_cache=cache) + await delete_cache_access_object(access_group_id=ag_id, user_api_key_cache=cache) cached = await cache.async_get_cache(key=f"access_group_id:{ag_id}") assert cached is None @@ -1047,8 +1047,8 @@ async def test_get_resources_from_access_groups( from litellm.proxy._types import LiteLLM_AccessGroupTable from litellm.proxy.auth.auth_checks import ( - _get_agent_ids_from_access_groups, - _get_models_from_access_groups, + get_agent_ids_from_access_groups, + get_models_from_access_groups, ) ag_table = LiteLLM_AccessGroupTable( @@ -1064,13 +1064,13 @@ async def test_get_resources_from_access_groups( return_value=ag_table, ): if resource_field == "access_model_names": - result = await _get_models_from_access_groups( + result = await get_models_from_access_groups( access_group_ids=[access_group_data["access_group_id"]], prisma_client=MagicMock(), user_api_key_cache=DualCache(), ) else: - result = await _get_agent_ids_from_access_groups( + result = await get_agent_ids_from_access_groups( access_group_ids=[access_group_data["access_group_id"]], prisma_client=MagicMock(), user_api_key_cache=DualCache(), @@ -1081,9 +1081,9 @@ async def test_get_resources_from_access_groups( @pytest.mark.asyncio async def test_get_models_from_access_groups_empty_ids(): """Test _get_models_from_access_groups returns empty list when access_group_ids is empty.""" - from litellm.proxy.auth.auth_checks import _get_models_from_access_groups + from litellm.proxy.auth.auth_checks import get_models_from_access_groups - result = await _get_models_from_access_groups(access_group_ids=[]) + result = await get_models_from_access_groups(access_group_ids=[]) assert result == [] @@ -1106,7 +1106,7 @@ async def test_can_team_access_model_via_access_group_ids(): ) with patch( - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new_callable=AsyncMock, return_value=["gpt-4"], ): @@ -1134,7 +1134,7 @@ async def test_can_team_access_model_access_group_ids_denied(): ) with patch( - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new_callable=AsyncMock, return_value=["claude-3"], ): @@ -1174,7 +1174,7 @@ async def test_can_key_call_model_via_access_group_ids(): ) with patch( - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new_callable=AsyncMock, return_value=["gpt-4"], ): @@ -1231,7 +1231,7 @@ async def test_key_access_group_grants_model_when_team_authorized(): """ from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token", @@ -1262,7 +1262,7 @@ async def test_key_access_group_grants_model_when_team_authorized(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1284,7 +1284,7 @@ async def test_key_access_group_grants_model_when_key_directly_authorized(): """ from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token-hashed", @@ -1316,7 +1316,7 @@ async def test_key_access_group_grants_model_when_key_directly_authorized(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1332,7 +1332,7 @@ async def test_key_access_group_grants_model_when_key_directly_authorized(): @pytest.mark.asyncio async def test_key_access_group_grants_model_when_key_has_no_groups(): """Key with no access_group_ids → False (early return, no DB read).""" - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token", @@ -1346,7 +1346,7 @@ async def test_key_access_group_grants_model_when_key_has_no_groups(): access_group_ids=["any-group"], ) assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1361,7 +1361,7 @@ async def test_key_access_group_grants_model_when_group_does_not_cover_model(): """Group authorizes the team but does not grant the requested model → False.""" from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token", @@ -1392,7 +1392,7 @@ async def test_key_access_group_grants_model_when_group_does_not_cover_model(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1415,7 +1415,7 @@ async def test_key_access_group_grants_model_when_group_authorizes_neither(): """ from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="team-a-token", @@ -1447,7 +1447,7 @@ async def test_key_access_group_grants_model_when_group_authorizes_neither(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-opus-4-5", valid_token=valid_token, team_object=team_object, @@ -1465,7 +1465,7 @@ async def test_key_access_group_grants_model_when_get_access_object_raises(): """Group lookup failure (404, network, etc.) is treated as no authorization.""" from unittest.mock import AsyncMock, patch - from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + from litellm.proxy.auth.auth_checks import key_access_group_grants_model valid_token = UserAPIKeyAuth( token="test-token", @@ -1490,7 +1490,7 @@ async def test_key_access_group_grants_model_when_get_access_object_raises(): p.start() try: assert ( - await _key_access_group_grants_model( + await key_access_group_grants_model( model="claude-haiku-4-5", valid_token=valid_token, team_object=team_object, @@ -1633,7 +1633,7 @@ def test_is_model_cost_zero_judges_an_alias_chain_by_the_deployment_its_entry_ro expected: Final = {"chain-entry": True, "local-free": False, "paid-gpt": False} order: Final = ("chain-entry", "local-free", "paid-gpt") if entry_first else ("local-free", "paid-gpt", "chain-entry") - verdicts: Final = {name: _is_model_cost_zero(model=name, llm_router=router) for name in order} + verdicts: Final = {name: is_model_cost_zero(model=name, llm_router=router) for name in order} assert verdicts == expected - assert {name: _is_model_cost_zero(model=name, llm_router=router) for name in order} == expected + assert {name: is_model_cost_zero(model=name, llm_router=router) for name in order} == expected diff --git a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py index e855d8cd346..e26845e5b2a 100644 --- a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py +++ b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py @@ -43,24 +43,24 @@ from litellm.proxy.auth.auth_checks import ( LITELLM_SESSION_TOKEN_PREFIX, ExperimentalUIJWTToken, _cache_management_object, - _can_object_call_model, + can_object_call_model, _can_object_call_vector_stores, _check_agent_access_group_model_access, - _check_end_user_budget, + check_end_user_budget, _check_team_member_budget, - _fetch_key_object_from_db_with_reconnect, + fetch_key_object_from_db_with_reconnect, _get_fuzzy_user_object, CallerTeamLoader, CallerUserLoader, _get_team_db_check, _log_budget_lookup_failure, _tag_max_budget_check, - _team_max_budget_check, - _team_member_max_budget_alert_check, - _virtual_key_max_budget_alert_check, + team_max_budget_check, + team_member_max_budget_alert_check, + virtual_key_max_budget_alert_check, _check_agent_caller_model_access, - _virtual_key_max_budget_check, - _virtual_key_soft_budget_check, + virtual_key_max_budget_check, + virtual_key_soft_budget_check, common_checks, get_key_object, get_user_object, @@ -343,7 +343,7 @@ def test_get_key_object_from_ui_hash_key_invalid(): ) def test_can_object_call_model_denials_return_forbidden(object_type, expected_error_type): with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="restricted-model", llm_router=None, models=["allowed-model"], @@ -481,7 +481,7 @@ async def test_enforce_key_access_teamless_all_team_models_passes(): the sentinel is present, regardless of team_id. Fails if someone adds a team_id guard to the pass branch.""" from litellm.proxy._types import SpecialModelNames - from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access + from litellm.proxy.auth.user_api_key_auth import enforce_key_and_fallback_model_access valid_token = UserAPIKeyAuth( api_key="sk-orphan", @@ -489,7 +489,7 @@ async def test_enforce_key_access_teamless_all_team_models_passes(): team_models=[], ) - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data={"model": "gpt-4o"}, route="/chat/completions", @@ -576,7 +576,7 @@ async def test_can_team_access_model_error_lists_direct_and_access_group_models( ) with patch( # test-quality-ok: access-group lookup has no dependency-injection seam - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new=AsyncMock(return_value=["group-model"]), ): assert await can_team_access_model("direct-model", team_object, None) is True @@ -669,7 +669,7 @@ async def test_fetch_key_object_from_db_bounds_in_flight_prisma_requests(): results: Final = await asyncio.gather( *( - _fetch_key_object_from_db_with_reconnect( + fetch_key_object_from_db_with_reconnect( hashed_token=f"hashed-token-{i}", prisma_client=prisma, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient parent_otel_span=None, @@ -730,7 +730,7 @@ async def test_fetch_key_object_from_db_fails_a_stalled_burst_within_the_deadlin results: Final = await asyncio.gather( *( - _fetch_key_object_from_db_with_reconnect( + fetch_key_object_from_db_with_reconnect( hashed_token=f"hashed-token-{i}", prisma_client=prisma, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient parent_otel_span=None, @@ -753,7 +753,7 @@ async def test_fetch_key_object_from_db_fails_a_stalled_burst_within_the_deadlin after: Final = await asyncio.wait_for( asyncio.gather( *( - _fetch_key_object_from_db_with_reconnect( + fetch_key_object_from_db_with_reconnect( hashed_token=f"after-{i}", prisma_client=recovered, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient parent_otel_span=None, @@ -2379,7 +2379,7 @@ async def test_key_and_team_grants_are_read_through_the_object_permission_cache( def test_can_object_call_model_with_alias(): """Test that can_object_call_model works with model aliases""" from litellm import Router - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model model = "[ip-approved] gpt-4o" llm_router = Router( @@ -2400,7 +2400,7 @@ def test_can_object_call_model_with_alias(): }, ) - result = _can_object_call_model( + result = can_object_call_model( model=model, llm_router=llm_router, models=["gpt-3.5-turbo"], @@ -2423,7 +2423,7 @@ def test_can_object_call_model_access_via_alias_only(): - The call should succeed because access is granted via the alias """ from litellm import Router - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model model = "my-fake-gpt" llm_router = Router( @@ -2445,7 +2445,7 @@ def test_can_object_call_model_access_via_alias_only(): ) # Key has access to the alias but NOT the underlying model - result = _can_object_call_model( + result = can_object_call_model( model=model, llm_router=llm_router, models=["my-fake-gpt"], # Only has access to alias, not "gpt-4" @@ -2460,9 +2460,9 @@ def test_can_object_call_model_access_via_alias_only(): def test_can_object_call_model_key_alias_to_allowed_target_is_allowed(): """A key alias whose target is on the key allowlist resolves like a team alias.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model - result = _can_object_call_model( + result = can_object_call_model( model="mistral-7b", llm_router=None, models=["gpt-4o-mini"], @@ -2477,10 +2477,10 @@ def test_can_object_call_model_key_alias_to_allowed_target_is_allowed(): def test_can_object_call_model_key_alias_to_disallowed_target_is_denied(): """A key alias whose target is outside the key allowlist stays denied.""" from litellm.proxy._types import ProxyErrorTypes, ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="mistral-7b", llm_router=None, models=["gpt-4o-mini"], @@ -2563,12 +2563,12 @@ async def test_can_key_call_model_honors_key_alias(): def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch): """The key alias rewrite precedes the global one at dispatch, so the key target is authorized.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["baz"], @@ -2580,7 +2580,7 @@ def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch ) with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["bar"], @@ -2594,12 +2594,12 @@ def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypatch): """A key alias on the globally rewritten name resolves the same way the request chain does.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["baz"], @@ -2613,12 +2613,12 @@ def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypat def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): """When a key alias fires on the globally rewritten name, only the final target is dispatched.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["bar"], @@ -2630,7 +2630,7 @@ def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["baz"], @@ -2644,10 +2644,10 @@ def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): def test_can_object_call_model_key_alias_name_alone_is_not_enough(): """A key that may call the alias name but not its target cannot call the alias.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="bar", llm_router=None, models=["bar"], @@ -2659,7 +2659,7 @@ def test_can_object_call_model_key_alias_name_alone_is_not_enough(): assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied assert ( - _can_object_call_model( + can_object_call_model( model="bar", llm_router=None, models=["baz"], @@ -2673,10 +2673,10 @@ def test_can_object_call_model_key_alias_name_alone_is_not_enough(): def test_can_object_call_model_team_alias_applies_before_key_alias(): """A key alias on the raw name loses to the team alias that rewrites it first at dispatch.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["bar"], @@ -2691,10 +2691,10 @@ def test_can_object_call_model_team_alias_applies_before_key_alias(): def test_can_object_call_model_key_alias_on_team_alias_target(): """A key alias on the team-rewritten name resolves like the dispatch chain does.""" - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model assert ( - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["baz"], @@ -2707,7 +2707,7 @@ def test_can_object_call_model_key_alias_on_team_alias_target(): ) with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="foo", llm_router=None, models=["bar"], @@ -2751,7 +2751,7 @@ async def test_can_user_call_model_honors_key_alias(): async def test_check_team_member_model_access_honors_key_alias(): """A key alias resolves against the member allowlist, not just the raw alias name.""" from litellm.proxy._types import LiteLLM_TeamMembership - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access membership = LiteLLM_TeamMembership( user_id="alice", @@ -2759,7 +2759,7 @@ async def test_check_team_member_model_access_honors_key_alias(): litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["gpt-4o-mini"]), ) - await _check_team_member_model_access( + await check_team_member_model_access( model="mistral-7b", team_object=LiteLLM_TeamTable(team_id="team-a"), valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), @@ -2773,7 +2773,7 @@ async def test_check_team_member_model_access_honors_key_alias(): ) with pytest.raises(ProxyException) as exc_info: - await _check_team_member_model_access( + await check_team_member_model_access( model="mistral-7b", team_object=LiteLLM_TeamTable(team_id="team-a"), valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), @@ -2799,7 +2799,7 @@ def test_can_object_call_model_access_via_underlying_model_only(): - The call should succeed because access is granted via the underlying model """ from litellm import Router - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model model = "my-fake-gpt" llm_router = Router( @@ -2821,7 +2821,7 @@ def test_can_object_call_model_access_via_underlying_model_only(): ) # Key has access to the underlying model but NOT the alias - result = _can_object_call_model( + result = can_object_call_model( model=model, llm_router=llm_router, models=["gpt-4"], # Only has access to underlying model, not "my-fake-gpt" @@ -2840,7 +2840,7 @@ def test_can_object_call_model_no_access_to_alias_or_underlying(): """ from litellm import Router from litellm.proxy._types import ProxyErrorTypes, ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model model = "my-fake-gpt" llm_router = Router( @@ -2863,7 +2863,7 @@ def test_can_object_call_model_no_access_to_alias_or_underlying(): # Key has access to neither the alias nor the underlying model with pytest.raises(ProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model=model, llm_router=llm_router, models=["gpt-3.5-turbo"], # Has access to different model entirely @@ -2887,7 +2887,7 @@ _DENIED_MESSAGE_TEMPLATE: Final = ( def test_can_object_call_model_denial_hides_allowlist_and_keeps_detail_on_exception(caplog): with caplog.at_level("DEBUG", logger="LiteLLM Proxy"): with pytest.raises(ModelAccessDeniedProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="anthropic-sonnet-4-5", llm_router=None, models=["internal-models"], @@ -2914,7 +2914,7 @@ async def test_access_group_fallback_grant_does_not_log_a_denial(caplog): with ( patch( # test-quality-ok: access-group lookup has no dependency-injection seam - "litellm.proxy.auth.auth_checks._get_models_from_access_groups", + "litellm.proxy.auth.auth_checks.get_models_from_access_groups", new=AsyncMock(return_value=["group-model"]), ), caplog.at_level("DEBUG", logger="LiteLLM Proxy"), @@ -2934,7 +2934,7 @@ async def test_access_group_fallback_grant_does_not_log_a_denial(caplog): ) def test_can_object_call_model_denial_same_client_message_for_every_object_type(object_type, expected_type): with pytest.raises(ModelAccessDeniedProxyException) as exc_info: - _can_object_call_model( + can_object_call_model( model="anthropic-sonnet-4-5", llm_router=None, models=["internal-models"], @@ -2964,7 +2964,7 @@ async def test_can_user_call_model_no_default_models_hides_policy_detail(): @pytest.mark.asyncio async def test_check_team_member_model_access_denied_hides_member_allowlist(): from litellm.proxy._types import LiteLLM_TeamMembership - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key membership = LiteLLM_TeamMembership( @@ -2980,7 +2980,7 @@ async def test_check_team_member_model_access_denied_hides_member_allowlist(): ) with pytest.raises(ModelAccessDeniedProxyException) as exc_info: - await _check_team_member_model_access( + await check_team_member_model_access( model="mock-vision", team_object=LiteLLM_TeamTable(team_id="team-a"), valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), @@ -3058,11 +3058,11 @@ def test_can_object_call_model_access_group_with_team_id(): model_info.access_groups for team-scoped DB models and allow access via group name. """ - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() - result = _can_object_call_model( + result = can_object_call_model( model="mock-fast-1", llm_router=router, models=["fast-models", "mock-power"], @@ -3079,12 +3079,12 @@ def test_can_object_call_model_access_group_without_team_id_fails(): This is the pre-fix behavior. """ from litellm.proxy._types import ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() with pytest.raises(ProxyException): - _can_object_call_model( + can_object_call_model( model="mock-fast-1", llm_router=router, models=["fast-models", "mock-power"], @@ -3098,11 +3098,11 @@ def test_can_object_call_model_literal_name_with_team_id(): Literal model name matching should still work when team_id is passed — no regression from adding team_id. """ - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() - result = _can_object_call_model( + result = can_object_call_model( model="mock-power", llm_router=router, models=["fast-models", "mock-power"], @@ -3118,12 +3118,12 @@ def test_can_object_call_model_denied_model_with_team_id(): still be denied even when team_id is passed. """ from litellm.proxy._types import ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() with pytest.raises(ProxyException): - _can_object_call_model( + can_object_call_model( model="mock-vision", llm_router=router, models=["fast-models", "mock-power"], @@ -3137,11 +3137,11 @@ def test_can_object_call_model_second_group_member_with_team_id(): Both models in the access group should be reachable, not just the first one. """ - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_team_scoped_router() - result = _can_object_call_model( + result = can_object_call_model( model="mock-fast-2", llm_router=router, models=["fast-models"], @@ -3164,7 +3164,7 @@ async def test_check_team_member_model_access_with_access_group(): LiteLLM_TeamTable, UserAPIKeyAuth, ) - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access router = _make_team_scoped_router() team = LiteLLM_TeamTable(team_id="team-a") @@ -3182,7 +3182,7 @@ async def test_check_team_member_model_access_with_access_group(): return_value=membership, ): # Should not raise — mock-fast-1 is in the fast-models group - await _check_team_member_model_access( + await check_team_member_model_access( model="mock-fast-1", team_object=team, valid_token=token, @@ -3206,7 +3206,7 @@ async def test_check_team_member_model_access_denied_model(): ProxyException, UserAPIKeyAuth, ) - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access router = _make_team_scoped_router() team = LiteLLM_TeamTable(team_id="team-a") @@ -3224,7 +3224,7 @@ async def test_check_team_member_model_access_denied_model(): return_value=membership, ): with pytest.raises(ProxyException) as exc_info: - await _check_team_member_model_access( + await check_team_member_model_access( model="mock-vision", team_object=team, valid_token=token, @@ -3249,7 +3249,7 @@ async def test_check_team_member_model_access_no_override_inherits_team(): LiteLLM_TeamTable, UserAPIKeyAuth, ) - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access router = _make_team_scoped_router() team = LiteLLM_TeamTable(team_id="team-a") @@ -3265,7 +3265,7 @@ async def test_check_team_member_model_access_no_override_inherits_team(): return_value=membership, ): # Should return without raising — no per-member restriction - await _check_team_member_model_access( + await check_team_member_model_access( model="mock-vision", team_object=team, valid_token=token, @@ -4278,7 +4278,7 @@ async def test_virtual_key_soft_budget_check_with_user_obj(): proxy_logging_obj = MockProxyLogging() - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -4326,7 +4326,7 @@ async def test_virtual_key_soft_budget_check_without_user_obj(): proxy_logging_obj = MockProxyLogging() - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=None, @@ -4370,7 +4370,7 @@ async def test_virtual_key_soft_budget_check_scenarios(spend, soft_budget, expec proxy_logging_obj = MockProxyLogging() - await _virtual_key_soft_budget_check( + await virtual_key_soft_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=None, @@ -4417,7 +4417,7 @@ async def test_virtual_key_max_budget_alert_check_with_user_obj(): proxy_logging_obj = MockProxyLogging() - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=user_obj, @@ -4465,7 +4465,7 @@ async def test_virtual_key_max_budget_alert_check_without_user_obj(): proxy_logging_obj = MockProxyLogging() - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=None, @@ -4503,7 +4503,7 @@ async def test_team_member_max_budget_alert_check_dispatches_only_at_configured_ async def budget_alerts(self, type, user_info): captured.append((type, user_info)) - _team_member_max_budget_alert_check( + team_member_max_budget_alert_check( team_id="team-1", team_alias="platform", team_metadata=team_metadata, @@ -4537,7 +4537,7 @@ async def test_team_member_max_budget_alert_check_drops_thresholds_outside_1_to_ async def budget_alerts(self, type, user_info): captured.append(user_info) - _team_member_max_budget_alert_check( + team_member_max_budget_alert_check( team_id="team-1", team_alias="platform", team_metadata={ @@ -4647,7 +4647,7 @@ async def test_virtual_key_max_budget_alert_check_scenarios(spend, max_budget, e proxy_logging_obj = MockProxyLogging() - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, user_obj=None, @@ -4690,7 +4690,7 @@ async def test_virtual_key_max_budget_alert_check_with_multi_threshold_map(): max_budget=None, ) - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=user_obj, @@ -4726,7 +4726,7 @@ async def test_virtual_key_max_budget_alert_check_old_path_no_map(): metadata={}, ) - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=None, @@ -4758,7 +4758,7 @@ async def test_virtual_key_max_budget_alert_check_old_path_below_threshold_no_al metadata={}, ) - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=None, @@ -4798,7 +4798,7 @@ async def test_virtual_key_max_budget_alert_check_global_fallback(): original = litellm.default_key_max_budget_alert_emails try: litellm.default_key_max_budget_alert_emails = global_config - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=None, @@ -4838,7 +4838,7 @@ async def test_virtual_key_max_budget_alert_check_per_key_merges_with_global(): original = litellm.default_key_max_budget_alert_emails try: litellm.default_key_max_budget_alert_emails = global_config - await _virtual_key_max_budget_alert_check( + await virtual_key_max_budget_alert_check( valid_token=valid_token, proxy_logging_obj=MockProxyLogging(), user_obj=None, @@ -4906,7 +4906,7 @@ async def test_custom_auth_common_checks_opt_in(): the pre-existing RPS guarantee for custom-auth hot paths. """ import litellm.proxy.proxy_server as _proxy_server_mod - from litellm.proxy.auth.user_api_key_auth import _run_centralized_common_checks + from litellm.proxy.auth.user_api_key_auth import run_centralized_common_checks valid_token = UserAPIKeyAuth(token="test-token", user_id="u1") mock_request = MagicMock() @@ -4934,7 +4934,7 @@ async def test_custom_auth_common_checks_opt_in(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_common: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=valid_token, request=mock_request, request_data={}, @@ -4955,7 +4955,7 @@ async def test_custom_auth_common_checks_opt_in(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_common: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=valid_token, request=mock_request, request_data={}, @@ -4995,7 +4995,7 @@ async def test_virtual_key_budget_check_reads_from_spend_counter(): with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -5027,7 +5027,7 @@ async def test_virtual_key_budget_check_fallback_no_counter(): with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -5093,7 +5093,7 @@ async def test_budget_exceeded_throttles_instead_of_blocking(monkeypatch): ) with _patched_spend(20.0): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5111,12 +5111,12 @@ async def test_budget_exceeded_throttles_instead_of_blocking(monkeypatch): async def test_budget_throttle_decision_cleared_before_caching(): """The request-scoped throttle decision must not persist into the key cache, otherwise it would re-apply (and compound) on every subsequent request.""" - from litellm.proxy.auth.auth_checks import _copy_user_api_key_auth_for_cache + from litellm.proxy.auth.auth_checks import copy_user_api_key_auth_for_cache valid_token = _over_budget_token(tpm_limit=1000, rpm_limit=100, metadata={"throttle_on_budget_exceeded": True}) valid_token.budget_throttle_pct = 0.1 - cached = _copy_user_api_key_auth_for_cache(user_api_key_obj=valid_token) + cached = copy_user_api_key_auth_for_cache(user_api_key_obj=valid_token) assert cached.budget_throttle_pct is None assert cached.tpm_limit == 1000 @@ -5132,7 +5132,7 @@ async def test_budget_exceeded_throttle_no_configured_limits(monkeypatch): with _patched_spend(20.0): with pytest.raises(litellm.BudgetExceededError): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5147,7 +5147,7 @@ async def test_budget_exceeded_not_opted_in_still_blocks(monkeypatch): with _patched_spend(20.0): with pytest.raises(litellm.BudgetExceededError): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5167,7 +5167,7 @@ async def test_budget_exceeded_invalid_percentage_blocks(monkeypatch, pct): with _patched_spend(20.0): with pytest.raises(litellm.BudgetExceededError): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5186,7 +5186,7 @@ async def test_under_budget_does_not_throttle(monkeypatch): ) with _patched_spend(5.0): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_budget_logging_obj(), ) @@ -5216,7 +5216,7 @@ async def test_team_budget_check_reads_from_spend_counter(): with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _team_max_budget_check( + await team_max_budget_check( team_object=team_object, valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, @@ -5243,7 +5243,7 @@ async def test_end_user_budget_check_reads_from_spend_counter(): with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _check_end_user_budget( + await check_end_user_budget( end_user_obj=end_user_object, route="/chat/completions", ) @@ -6311,7 +6311,7 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): from unittest.mock import AsyncMock, MagicMock from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object + from litellm.proxy.auth.auth_checks import cache_team_object base_team_row = { "team_id": "team-1234", @@ -6328,7 +6328,7 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): logging_obj = MagicMock() logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() - await _cache_team_object( + await cache_team_object( team_id="team-1234", team_table=team_table, user_api_key_cache=cache, @@ -6365,7 +6365,7 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): logging_obj2 = MagicMock() logging_obj2.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() - await _cache_team_object( + await cache_team_object( team_id="team-no-alias", team_table=aliasless, user_api_key_cache=cache2, @@ -6427,7 +6427,7 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): """ from litellm.caching.dual_cache import DualCache from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object + from litellm.proxy.auth.auth_checks import cache_team_object, get_team_object team_id = "team-lit-4391" shared_redis = _SharedFakeRedis() @@ -6439,7 +6439,7 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): ) prisma_client = MagicMock() - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]), user_api_key_cache=user_api_key_cache, @@ -6454,7 +6454,7 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): ) assert primed is not None and primed.models == ["model-a"] - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a", "model-b"]), user_api_key_cache=user_api_key_cache, @@ -6514,7 +6514,7 @@ async def test_warm_team_object_reads_issue_no_redis_ops_lit_5944(): """ from litellm.caching.dual_cache import DualCache from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object + from litellm.proxy.auth.auth_checks import cache_team_object, get_team_object team_id = "team-lit-5944" counting_redis = _CountingFakeRedis() @@ -6526,7 +6526,7 @@ async def test_warm_team_object_reads_issue_no_redis_ops_lit_5944(): ) prisma_client = MagicMock() - await _cache_team_object( + await cache_team_object( team_id=team_id, team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]), user_api_key_cache=user_api_key_cache, @@ -6560,7 +6560,7 @@ async def test_cache_team_object_tolerates_cache_invalidation_failures(): a 500. The authoritative team_id-keyed write must still happen. """ from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object + from litellm.proxy.auth.auth_checks import cache_team_object cache = MagicMock() cache.async_set_cache = AsyncMock() @@ -6568,7 +6568,7 @@ async def test_cache_team_object_tolerates_cache_invalidation_failures(): logging_obj = MagicMock() logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock(side_effect=Exception("redis down")) - await _cache_team_object( + await cache_team_object( team_id="team-cache-outage", team_table=LiteLLM_TeamTableCachedObj( team_id="team-cache-outage", @@ -6711,7 +6711,7 @@ async def test_virtual_key_max_budget_error_names_the_key(): new=AsyncMock(return_value=25.0), ): with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -6737,7 +6737,7 @@ async def test_virtual_key_max_budget_not_exceeded_does_not_raise(): "litellm.proxy.proxy_server.get_current_spend", new=AsyncMock(return_value=1.0), ): - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -6990,7 +6990,7 @@ async def test_common_checks_personal_user_budget_blocks_in_gather(): async def _common_checks_for_over_budget_personal_key(*, model: str) -> bool: from litellm import Router - from litellm.proxy.auth.auth_checks import _is_model_cost_zero, common_checks + from litellm.proxy.auth.auth_checks import is_model_cost_zero, common_checks llm_router: Final = Router( model_list=[ @@ -7030,7 +7030,7 @@ async def _common_checks_for_over_budget_personal_key(*, model: str) -> bool: proxy_logging_obj=proxy_logging_obj, valid_token=token, request=MagicMock(spec=Request), - skip_budget_checks=_is_model_cost_zero(model=model, llm_router=llm_router), + skip_budget_checks=is_model_cost_zero(model=model, llm_router=llm_router), ) await asyncio.sleep(0) return result @@ -7388,7 +7388,7 @@ async def test_organization_budget_check_carries_org_state_on_the_token(): (Prometheus org budget gauges) reads it from request metadata instead of calling get_org_object again.""" from litellm.proxy._types import LiteLLM_OrganizationTable - from litellm.proxy.auth.auth_checks import _organization_max_budget_check + from litellm.proxy.auth.auth_checks import organization_max_budget_check from litellm.types.proxy.carried_budget_state import OrgBudgetSnapshot org_table = LiteLLM_OrganizationTable( @@ -7406,7 +7406,7 @@ async def test_organization_budget_check_carries_org_state_on_the_token(): key="org_id:o1:with_budget", value=org_table, model_type=LiteLLM_OrganizationTable ) - await _organization_max_budget_check( + await organization_max_budget_check( valid_token=token, team_object=None, prisma_client=MagicMock(), @@ -7540,7 +7540,7 @@ async def test_organization_zero_max_budget_is_enforced(max_budget, spend, expec spend without limit. """ from litellm.proxy._types import LiteLLM_OrganizationTable - from litellm.proxy.auth.auth_checks import _organization_max_budget_check + from litellm.proxy.auth.auth_checks import organization_max_budget_check org_table = LiteLLM_OrganizationTable( organization_id="o1", @@ -7568,7 +7568,7 @@ async def test_organization_zero_max_budget_is_enforced(max_budget, spend, expec ): if expect_blocked: with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _organization_max_budget_check( + await organization_max_budget_check( valid_token=token, team_object=None, prisma_client=MagicMock(), @@ -7577,7 +7577,7 @@ async def test_organization_zero_max_budget_is_enforced(max_budget, spend, expec ) assert exc_info.value.max_budget == max_budget else: - await _organization_max_budget_check( + await organization_max_budget_check( valid_token=token, team_object=None, prisma_client=MagicMock(), @@ -8692,11 +8692,11 @@ def _restricted_member_check_deps() -> dict[str, object]: @pytest.mark.asyncio async def test_check_team_member_model_access_fails_closed_when_the_membership_read_hits_a_db_outage(): - from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_checks import check_team_member_model_access from litellm.proxy.auth.auth_exception_handler import _as_proxy_exception with pytest.raises(httpx.ConnectError) as raised: - await _check_team_member_model_access( + await check_team_member_model_access( model="claude-sonnet-5", llm_router=None, **_restricted_member_check_deps() ) @@ -9105,7 +9105,7 @@ def test_is_user_proxy_admin_rejects_view_only_admin(): """This predicate skips `non_proxy_admin_allowed_routes_check` entirely, so an Admin Viewer answering True here would gain every write route. Read parity for that role belongs in the route checks, never here.""" - from litellm.proxy.auth.auth_checks import _is_user_proxy_admin + from litellm.proxy.auth.auth_checks import is_user_proxy_admin viewer = LiteLLM_UserTable( user_id="viewer_user", @@ -9118,9 +9118,9 @@ def test_is_user_proxy_admin_rejects_view_only_admin(): user_role=LitellmUserRoles.PROXY_ADMIN.value, ) - assert _is_user_proxy_admin(user_obj=viewer) is False - assert _is_user_proxy_admin(user_obj=admin) is True - assert _is_user_proxy_admin(user_obj=None) is False + assert is_user_proxy_admin(user_obj=viewer) is False + assert is_user_proxy_admin(user_obj=admin) is True + assert is_user_proxy_admin(user_obj=None) is False def _make_wildcard_access_group_router(): @@ -9156,12 +9156,12 @@ def test_can_object_call_model_access_group_wildcard_accepts_bare_model_name(): pattern router's raw regex and skipped the `{provider}/{model}` retry that both routing and the direct-wildcard grant already perform. """ - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_wildcard_access_group_router() assert ( - _can_object_call_model( + can_object_call_model( model="gpt-4o", llm_router=router, models=["default-models"], @@ -9172,12 +9172,12 @@ def test_can_object_call_model_access_group_wildcard_accepts_bare_model_name(): def test_can_object_call_model_access_group_wildcard_accepts_prefixed_model_name(): - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_wildcard_access_group_router() assert ( - _can_object_call_model( + can_object_call_model( model="openai/gpt-4o", llm_router=router, models=["default-models"], @@ -9197,12 +9197,12 @@ def test_can_object_call_model_access_group_wildcard_accepts_prefixed_model_name def test_can_object_call_model_access_group_wildcard_does_not_over_grant(model): """The bare-name retry must not turn an access group into a blanket grant.""" from litellm.proxy._types import ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = _make_wildcard_access_group_router() with pytest.raises(ProxyException): - _can_object_call_model( + can_object_call_model( model=model, llm_router=router, models=["default-models"], @@ -9217,7 +9217,7 @@ def test_can_object_call_model_access_group_rejects_unconsumed_namespace(): """ from litellm import Router from litellm.proxy._types import ProxyException - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = Router( model_list=[ @@ -9233,7 +9233,7 @@ def test_can_object_call_model_access_group_rejects_unconsumed_namespace(): ) assert ( - _can_object_call_model( + can_object_call_model( model="anthropic.claude-3-5-sonnet-20240620-v1:0", llm_router=router, models=["bedrock-models"], @@ -9243,7 +9243,7 @@ def test_can_object_call_model_access_group_rejects_unconsumed_namespace(): ) with pytest.raises(ProxyException): - _can_object_call_model( + can_object_call_model( model="bedrockz/anthropic.claude-3-5-sonnet-20240620-v1:0", llm_router=router, models=["bedrock-models"], @@ -9258,7 +9258,7 @@ def test_can_object_call_model_team_scoped_wildcard_accepts_bare_model_name(): index that needed the same `{provider}/{model}` retry. """ from litellm import Router - from litellm.proxy.auth.auth_checks import _can_object_call_model + from litellm.proxy.auth.auth_checks import can_object_call_model router = Router( model_list=[ @@ -9277,7 +9277,7 @@ def test_can_object_call_model_team_scoped_wildcard_accepts_bare_model_name(): for model in ("gpt-4o", "openai/gpt-4o"): assert ( - _can_object_call_model( + can_object_call_model( model=model, llm_router=router, models=["team-models"], @@ -10211,7 +10211,7 @@ async def test_delete_cache_key_object_is_best_effort_when_the_cache_backend_fai import logging from unittest.mock import AsyncMock, MagicMock - from litellm.proxy.auth.auth_checks import _delete_cache_key_object + from litellm.proxy.auth.auth_checks import delete_cache_key_object hashed_token = "a" * 64 caplog.set_level(logging.WARNING, logger="LiteLLM Proxy") @@ -10223,7 +10223,7 @@ async def test_delete_cache_key_object_is_best_effort_when_the_cache_backend_fai side_effect=Exception("No permissions to access a key") ) - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=failing_cache, proxy_logging_obj=failing_logging_obj, @@ -10241,7 +10241,7 @@ async def test_delete_cache_key_object_is_best_effort_when_the_cache_backend_fai healthy_logging_obj = MagicMock() healthy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token=hashed_token, user_api_key_cache=healthy_cache, proxy_logging_obj=healthy_logging_obj, @@ -10272,7 +10272,7 @@ async def _run_key_budget_check(key_name: str) -> str: max_budget=1.0, ) with pytest.raises(litellm.BudgetExceededError, match="Budget has been exceeded") as exc_info: - await _virtual_key_max_budget_check( + await virtual_key_max_budget_check( valid_token=valid_token, proxy_logging_obj=_BudgetAlertRecorder(), ) @@ -10920,7 +10920,7 @@ async def test_agent_key_without_an_echoed_caller_keeps_its_own_models(): def test_can_object_call_model_allows_listed_model_for_key(): - result: Final = _can_object_call_model( + result: Final = can_object_call_model( model="allowed-model", llm_router=None, models=["allowed-model"], @@ -11113,7 +11113,7 @@ async def test_authoritative_group_grants_propagate_policy_outages( from fastapi import HTTPException from litellm.proxy import proxy_server - from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups + from litellm.proxy.auth.auth_checks import get_agent_ids_from_access_groups from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache database: Final = MagicMock() @@ -11125,6 +11125,6 @@ async def test_authoritative_group_grants_propagate_policy_outages( monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache()) if strict: with pytest.raises(HTTPException): - await _get_agent_ids_from_access_groups(["group"], check_db_only=True) + await get_agent_ids_from_access_groups(["group"], check_db_only=True) else: - assert await _get_agent_ids_from_access_groups(["group"]) == [] + assert await get_agent_ids_from_access_groups(["group"]) == [] diff --git a/tests/unit/proxy/auth/test_auth_exception_handler.py b/tests/unit/proxy/auth/test_auth_exception_handler.py index 521cbd8daad..dd6e1f57fbc 100644 --- a/tests/unit/proxy/auth/test_auth_exception_handler.py +++ b/tests/unit/proxy/auth/test_auth_exception_handler.py @@ -68,7 +68,7 @@ async def test_handle_authentication_error_db_unavailable_connectivity(db_error) "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": True}, ): - result = await handler._handle_authentication_error( + result = await handler.handle_authentication_error( db_error, mock_request, {}, @@ -113,7 +113,7 @@ async def test_handle_authentication_error_permanent_fault_gets_no_fallback_iden {"allow_requests_on_db_unavailable": True}, ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( prisma_error, mock_request, {}, @@ -148,7 +148,7 @@ async def test_handle_authentication_error_permanent_fault_503_is_not_worded_as_ "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False} ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error(prisma_error, MagicMock(), {}, "/test", None, "test-key") + await handler.handle_authentication_error(prisma_error, MagicMock(), {}, "/test", None, "test-key") assert exc_info.value.code == str(status.HTTP_503_SERVICE_UNAVAILABLE) assert exc_info.value.type == ProxyErrorTypes.no_db_connection @@ -176,7 +176,7 @@ async def test_handle_authentication_error_transport_error_raised_over_a_permane "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False} ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error(transport_over_fault, MagicMock(), {}, "/test", None, "k") + await handler.handle_authentication_error(transport_over_fault, MagicMock(), {}, "/test", None, "k") assert exc_info.value.code == str(status.HTTP_503_SERVICE_UNAVAILABLE) assert "temporarily unreachable" not in exc_info.value.message @@ -202,7 +202,7 @@ async def test_handle_authentication_error_transient_outage_503_keeps_retry_word "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False} ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error(db_error, MagicMock(), {}, "/test", None, "test-key") + await handler.handle_authentication_error(db_error, MagicMock(), {}, "/test", None, "test-key") assert exc_info.value.code == str(status.HTTP_503_SERVICE_UNAVAILABLE) assert exc_info.value.message == ( @@ -249,7 +249,7 @@ async def test_handle_authentication_error_data_layer_errors_do_not_fall_back( {"allow_requests_on_db_unavailable": True}, ): with pytest.raises(ProxyException): - await handler._handle_authentication_error( + await handler.handle_authentication_error( prisma_error, mock_request, {}, @@ -300,7 +300,7 @@ async def test_handle_authentication_error_db_infra_error_returns_503(db_error): ), ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( db_error, MagicMock(), {}, @@ -357,7 +357,7 @@ async def test_handle_authentication_error_prisma_engine_teardown_returns_503(): ), ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( teardown_error, MagicMock(), {}, @@ -407,7 +407,7 @@ async def test_handle_authentication_error_genuine_auth_failure_stays_401(auth_e ), ): with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( auth_error, MagicMock(), {}, @@ -438,7 +438,7 @@ async def test_handle_authentication_error_budget_exceeded(): ) with pytest.raises(ProxyException) as exc_info: - await handler._handle_authentication_error( + await handler.handle_authentication_error( budget_error, mock_request, mock_request_data, @@ -476,7 +476,7 @@ async def test_route_passed_to_post_call_failure_hook(): {"allow_requests_on_db_unavailable": False}, ): try: - await handler._handle_authentication_error( + await handler.handle_authentication_error( PrismaError(), mock_request, mock_request_data, @@ -507,7 +507,7 @@ async def test_dynamic_route_normalized_on_auth_failure(): ), pytest.raises(ProxyException), ): - await handler._handle_authentication_error( + await handler.handle_authentication_error( HTTPException(status_code=401, detail="Authentication Error, Invalid proxy server token passed"), MagicMock(), {}, @@ -567,7 +567,7 @@ async def test_resolved_identity_exported_on_auth_failure(): ), ): with pytest.raises(ProxyException): - await handler._handle_authentication_error( + await handler.handle_authentication_error( expired_key_error, MagicMock(), {"model": "gpt-4o"}, @@ -666,7 +666,7 @@ async def test_expired_key_error_log_names_the_key_owner( verbose_proxy_logger.propagate = True try: with caplog.at_level("ERROR", logger="LiteLLM Proxy"), pytest.raises(ProxyException): - await handler._handle_authentication_error( + await handler.handle_authentication_error( expired_key_error, MagicMock(), {"model": "gpt-4o"}, @@ -709,7 +709,7 @@ async def test_auth_failure_without_resolved_identity_still_logs(): ), ): with pytest.raises(ProxyException): - await handler._handle_authentication_error( + await handler.handle_authentication_error( ProxyException( message="Invalid API key", type=ProxyErrorTypes.auth_error, @@ -1066,7 +1066,7 @@ async def test_handle_authentication_error_traceback_only_for_unexpected_errors( raise auth_error except (ProxyException, ValueError, HTTPException) as caught: with caplog.at_level(expect_level, logger="LiteLLM Proxy"), pytest.raises((ProxyException, HTTPException)): - await handler._handle_authentication_error( + await handler.handle_authentication_error( caught, MagicMock(), {}, @@ -1139,7 +1139,7 @@ async def test_handle_authentication_error_keeps_internal_message_on_model_acces caplog.at_level("WARNING", logger="LiteLLM Proxy"), pytest.raises(ModelAccessDeniedProxyException) as exc_info, ): - await handler._handle_authentication_error(denial, MagicMock(), {}, "/v1/chat/completions", None, "sk-bad-key") + await handler.handle_authentication_error(denial, MagicMock(), {}, "/v1/chat/completions", None, "sk-bad-key") assert exc_info.value.code == str(status.HTTP_403_FORBIDDEN) assert "internal-models" not in str(exc_info.value.message) diff --git a/tests/unit/proxy/auth/test_auth_utils.py b/tests/unit/proxy/auth/test_auth_utils.py index 08757534059..130d29a9560 100644 --- a/tests/unit/proxy/auth/test_auth_utils.py +++ b/tests/unit/proxy/auth/test_auth_utils.py @@ -706,7 +706,7 @@ def _cache_prediction_auth_app( from litellm.caching.dual_cache import DualCache from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LitellmUserRoles, ProxyException from litellm.proxy.auth import auth_checks - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.management_endpoints import prompt_cache_prediction as endpoint from litellm.proxy.utils import InternalUsageCache, ProxyLogging @@ -729,7 +729,7 @@ def _cache_prediction_auth_app( ) return token - monkeypatch.setattr(auth, "_user_api_key_auth_builder", authenticate) + monkeypatch.setattr(auth, "user_api_key_auth_builder", authenticate) monkeypatch.setattr(auth, "get_user_object", AsyncMock(return_value=user)) team = LiteLLM_TeamTableCachedObj(team_id=team_id, models=token.team_models) if team_id else None monkeypatch.setattr(auth, "get_team_object", AsyncMock(return_value=team)) @@ -744,7 +744,7 @@ def _cache_prediction_auth_app( monkeypatch.setattr(proxy_server, "prisma_client", None) monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache()) logging = ProxyLogging(user_api_key_cache=DualCache()) - logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + logging.proxy_hook_mapping["parallel_request_limiter"] = PROXY_MaxParallelRequestsHandler_v3( InternalUsageCache(dual_cache=DualCache()) ) monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging) diff --git a/tests/unit/proxy/auth/test_custom_auth_end_user_budget.py b/tests/unit/proxy/auth/test_custom_auth_end_user_budget.py index e83c5cf8419..0c3487255f6 100644 --- a/tests/unit/proxy/auth/test_custom_auth_end_user_budget.py +++ b/tests/unit/proxy/auth/test_custom_auth_end_user_budget.py @@ -151,11 +151,11 @@ async def test_custom_auth_defers_end_user_budget_to_common_checks_when_enabled( return_value=end_user_obj, ), patch( - "litellm.proxy.auth.user_api_key_auth._check_end_user_budget", + "litellm.proxy.auth.user_api_key_auth.check_end_user_budget", new_callable=AsyncMock, ) as mock_check, patch( - "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", + "litellm.proxy.auth.user_api_key_auth.enforce_key_and_fallback_model_access", new_callable=AsyncMock, ), patch( diff --git a/tests/unit/proxy/auth/test_default_end_user_budget_simple.py b/tests/unit/proxy/auth/test_default_end_user_budget_simple.py index edd0409343a..c58ea4c19c3 100644 --- a/tests/unit/proxy/auth/test_default_end_user_budget_simple.py +++ b/tests/unit/proxy/auth/test_default_end_user_budget_simple.py @@ -137,7 +137,7 @@ async def test_budget_enforcement_blocks_over_budget_users(): Note: Budget enforcement happens in common_checks() via _check_end_user_budget(), not in get_end_user_object(). get_end_user_object only fetches the user data. """ - from litellm.proxy.auth.auth_checks import _check_end_user_budget + from litellm.proxy.auth.auth_checks import check_end_user_budget end_user_id = f"test_user_{uuid.uuid4().hex}" default_budget_id = str(uuid.uuid4()) @@ -187,7 +187,7 @@ async def test_budget_enforcement_blocks_over_budget_users(): # Now test budget enforcement separately via _check_end_user_budget with pytest.raises(litellm.BudgetExceededError) as exc_info: - await _check_end_user_budget( + await check_end_user_budget( end_user_obj=result, route="/chat/completions", ) diff --git a/tests/unit/proxy/auth/test_login_utils.py b/tests/unit/proxy/auth/test_login_utils.py index e095af20b9c..1fb648fe5e6 100644 --- a/tests/unit/proxy/auth/test_login_utils.py +++ b/tests/unit/proxy/auth/test_login_utils.py @@ -639,7 +639,7 @@ class TestEncodeUiSessionJwt: from unittest.mock import MagicMock from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( - _user_id_from_session_cookie, + user_id_from_session_cookie, ) from litellm.proxy.auth.login_utils import encode_ui_session_jwt @@ -649,7 +649,7 @@ class TestEncodeUiSessionJwt: request = MagicMock() request.cookies = {"token": token} with patch("litellm.proxy.proxy_server.master_key", "sk-master-for-tests"): - assert _user_id_from_session_cookie(request) == "cornell-user" + assert user_id_from_session_cookie(request) == "cornell-user" def _throttle( diff --git a/tests/unit/proxy/auth/test_object_permission_loading.py b/tests/unit/proxy/auth/test_object_permission_loading.py index 8db4e210107..867ce9c998f 100644 --- a/tests/unit/proxy/auth/test_object_permission_loading.py +++ b/tests/unit/proxy/auth/test_object_permission_loading.py @@ -53,7 +53,7 @@ async def test_get_key_object_loads_object_permission(): "litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=mock_object_permission), ), - patch("litellm.proxy.auth.auth_checks._cache_key_object", AsyncMock()), + patch("litellm.proxy.auth.auth_checks.cache_key_object", AsyncMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), ): result = await get_key_object( @@ -94,7 +94,7 @@ async def test_get_key_object_no_permission_id(): mock_proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() with ( - patch("litellm.proxy.auth.auth_checks._cache_key_object", AsyncMock()), + patch("litellm.proxy.auth.auth_checks.cache_key_object", AsyncMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), ): result = await get_key_object( @@ -147,7 +147,7 @@ async def test_get_team_object_loads_object_permission(): "litellm.proxy.auth.auth_checks.get_object_permission", AsyncMock(return_value=mock_object_permission), ), - patch("litellm.proxy.auth.auth_checks._cache_team_object", AsyncMock()), + patch("litellm.proxy.auth.auth_checks.cache_team_object", AsyncMock()), patch("litellm.proxy.auth.auth_checks._should_check_db", return_value=True), patch("litellm.proxy.auth.auth_checks._update_last_db_access_time"), patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), diff --git a/tests/unit/proxy/auth/test_proxy_routes.py b/tests/unit/proxy/auth/test_proxy_routes.py index 129a93ea08d..24ef2a0e150 100644 --- a/tests/unit/proxy/auth/test_proxy_routes.py +++ b/tests/unit/proxy/auth/test_proxy_routes.py @@ -246,9 +246,9 @@ def _is_assistants(req): def _metadata_var_name(req): - from litellm.proxy.litellm_pre_call_utils import _get_metadata_variable_name + from litellm.proxy.litellm_pre_call_utils import get_metadata_variable_name - return _get_metadata_variable_name(req) + return get_metadata_variable_name(req) def _vector_store_id_in_path(req): diff --git a/tests/unit/proxy/auth/test_route_checks.py b/tests/unit/proxy/auth/test_route_checks.py index 36799b73876..540b66ddaba 100644 --- a/tests/unit/proxy/auth/test_route_checks.py +++ b/tests/unit/proxy/auth/test_route_checks.py @@ -14,7 +14,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import _is_api_route_allowed -from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin +from litellm.proxy.auth.auth_checks_organization import user_is_org_admin from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import router as llm_passthrough_router @@ -2946,7 +2946,7 @@ def test_available_roles_accessible_to_non_admin_users(user_role): ) -# ── _user_is_org_admin tests ────────────────────────────────────────────────── +# ── user_is_org_admin tests ────────────────────────────────────────────────── def _make_org_admin_user(org_id: str) -> LiteLLM_UserTable: @@ -2967,25 +2967,25 @@ def _make_org_admin_user(org_id: str) -> LiteLLM_UserTable: def test_user_is_org_admin_with_organizations_list(): """Org admin can be identified via the `organizations` list field (used by /user/new).""" user_obj = _make_org_admin_user("org-1") - assert _user_is_org_admin({"organizations": ["org-1"]}, user_obj) is True + assert user_is_org_admin({"organizations": ["org-1"]}, user_obj) is True def test_user_is_org_admin_with_singular_organization_id(): """Backward-compat: org admin can still be identified via singular `organization_id`.""" user_obj = _make_org_admin_user("org-1") - assert _user_is_org_admin({"organization_id": "org-1"}, user_obj) is True + assert user_is_org_admin({"organization_id": "org-1"}, user_obj) is True def test_user_is_org_admin_organizations_list_wrong_org(): """Non-member of the requested org is not considered an org admin for it.""" user_obj = _make_org_admin_user("org-2") - assert _user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False + assert user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False def test_user_is_org_admin_no_org_fields(): """Returns False when neither `organization_id` nor `organizations` is in the request.""" user_obj = _make_org_admin_user("org-1") - assert _user_is_org_admin({}, user_obj) is False + assert user_is_org_admin({}, user_obj) is False def test_non_org_admin_with_organizations_list(): @@ -3002,13 +3002,13 @@ def test_non_org_admin_with_organizations_list(): user_role=LitellmUserRoles.INTERNAL_USER.value, organization_memberships=[membership], ) - assert _user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False + assert user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False def test_org_admin_cannot_escalate_to_other_org(): """Regression: admin of org-A requesting [org-A, org-B] must be rejected.""" user_obj = _make_org_admin_user("org-A") - assert _user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is False + assert user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is False def test_org_admin_of_multiple_orgs_can_operate_on_both(): @@ -3034,7 +3034,7 @@ def test_org_admin_of_multiple_orgs_can_operate_on_both(): user_role=LitellmUserRoles.INTERNAL_USER.value, organization_memberships=memberships, ) - assert _user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is True + assert user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is True # ── LIT-4221: /team/update org-context resolution from team_id ──────────────── @@ -3742,7 +3742,7 @@ def test_organization_daily_activity_not_granted_by_org_admin_request_data_branc self_managed_routes entry is load-bearing rather than redundant. Query params do reach request_data, so the reason is not body-vs-query: it - is the key name. _user_is_org_admin reads ``organization_id`` (singular) and + is the key name. user_is_org_admin reads ``organization_id`` (singular) and ``organizations``, while this endpoint's filter is ``organization_ids`` (plural), and the dashboard's first page load sends no organization filter at all. Both shapes are pinned below because renaming the query param would @@ -3764,11 +3764,11 @@ def test_organization_daily_activity_not_granted_by_org_admin_request_data_branc ) # The dashboard's default page load: no organization filter at all. - assert not _user_is_org_admin(request_data={}, user_object=user_obj) + assert not user_is_org_admin(request_data={}, user_object=user_obj) # The filtered load, naming an org this user really does administer. - assert not _user_is_org_admin(request_data={"organization_ids": "org-a"}, user_object=user_obj) + assert not user_is_org_admin(request_data={"organization_ids": "org-a"}, user_object=user_obj) # The key name the helper would have had to see to grant it. - assert _user_is_org_admin(request_data={"organization_id": "org-a"}, user_object=user_obj) + assert user_is_org_admin(request_data={"organization_id": "org-a"}, user_object=user_obj) assert not RouteChecks.check_route_access( route="/organization/daily/activity", allowed_routes=LiteLLMRoutes.org_admin_only_routes.value, diff --git a/tests/unit/proxy/auth/test_router_override_fallback_auth.py b/tests/unit/proxy/auth/test_router_override_fallback_auth.py index c34614a93ed..d35a5851659 100644 --- a/tests/unit/proxy/auth/test_router_override_fallback_auth.py +++ b/tests/unit/proxy/auth/test_router_override_fallback_auth.py @@ -12,7 +12,7 @@ import pytest from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.auth_utils import fallback_target_model_name, iter_request_fallback_targets -from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access +from litellm.proxy.auth.user_api_key_auth import enforce_key_and_fallback_model_access def _fallback_model_names(fallbacks): @@ -100,7 +100,7 @@ async def test_router_override_fallbacks_validated_against_key_allowlist(): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -151,7 +151,7 @@ async def test_router_override_all_fallback_fields_validated(fallback_field): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -198,7 +198,7 @@ async def test_top_level_fallback_fields_validated(fallback_field): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -244,7 +244,7 @@ async def test_nested_deployment_fallback_inner_model_validated(): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -289,7 +289,7 @@ async def test_model_less_fallback_dict_is_skipped_never_passed_as_none(): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", @@ -327,7 +327,7 @@ async def test_router_override_without_fallbacks_does_not_break_auth(): new=AsyncMock(), ), ): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=valid_token, request_data=request_data, route="/v1/chat/completions", diff --git a/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py b/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py index be7b438a442..c2851f608f0 100644 --- a/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py +++ b/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py @@ -12,7 +12,7 @@ import copy import pytest import litellm -from litellm.proxy.auth.auth_checks import _is_model_cost_zero +from litellm.proxy.auth.auth_checks import is_model_cost_zero from litellm.router import Router @@ -88,7 +88,7 @@ class TestUnmappedModelBudgetEnforcement: }, ] ) - result = _is_model_cost_zero(model="custom-model", llm_router=router) + result = is_model_cost_zero(model="custom-model", llm_router=router) assert result is False, "Unmapped model should enforce budget (return False), not bypass it (return True)" def test_explicitly_free_model_bypasses_budget(self): @@ -111,7 +111,7 @@ class TestUnmappedModelBudgetEnforcement: }, ] ) - result = _is_model_cost_zero(model="free-model", llm_router=router) + result = is_model_cost_zero(model="free-model", llm_router=router) assert result is True, "Explicitly free model should bypass budget (return True)" def test_known_paid_model_enforces_budget(self): @@ -127,7 +127,7 @@ class TestUnmappedModelBudgetEnforcement: }, ] ) - result = _is_model_cost_zero(model="paid-model", llm_router=router) + result = is_model_cost_zero(model="paid-model", llm_router=router) assert result is False, "Known paid model should enforce budget (return False)" def test_unmapped_model_with_litellm_params_pricing(self): @@ -145,7 +145,7 @@ class TestUnmappedModelBudgetEnforcement: }, ] ) - result = _is_model_cost_zero(model="free-via-params", llm_router=router) + result = is_model_cost_zero(model="free-via-params", llm_router=router) assert result is True, "Model with explicit cost=0 in litellm_params should bypass budget" def test_cache_invalidates_on_in_place_pricing_update(self): @@ -177,7 +177,7 @@ class TestUnmappedModelBudgetEnforcement: ] ) # Warm the cache as zero-cost. - assert _is_model_cost_zero(model="ramping-model", llm_router=router) is True + assert is_model_cost_zero(model="ramping-model", llm_router=router) is True assert router._zero_cost_cache.get("ramping-model") is True # In-place pricing update: same deployment count, same router id, @@ -203,7 +203,7 @@ class TestUnmappedModelBudgetEnforcement: # Cache must have been cleared by ``_invalidate_model_group_info_cache``. assert router._zero_cost_cache == {} # Subsequent call sees the new pricing and enforces budget. - assert _is_model_cost_zero(model="ramping-model", llm_router=router) is False + assert is_model_cost_zero(model="ramping-model", llm_router=router) is False def test_strategy_router_alias_with_zero_pricing_enforces_budget(self): """An auto-router alias is never the deployment that gets called or @@ -231,7 +231,7 @@ class TestUnmappedModelBudgetEnforcement: ) assert "input_cost_per_token" not in litellm.model_cost.get("alias-id", {}) - assert _is_model_cost_zero(model="smart-router", llm_router=router) is False + assert is_model_cost_zero(model="smart-router", llm_router=router) is False def test_model_group_alias_to_free_model_bypasses_budget(self): """A zero-cost group reached through model_group_alias bypasses budget, like its own name. @@ -255,8 +255,8 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"free-model-alias": "free-model"}, ) - assert _is_model_cost_zero(model="free-model", llm_router=router) is True - assert _is_model_cost_zero(model="free-model-alias", llm_router=router) is True, ( + assert is_model_cost_zero(model="free-model", llm_router=router) is True + assert is_model_cost_zero(model="free-model-alias", llm_router=router) is True, ( "An alias pointing at an explicitly-zero-cost group must be read as free, like its own name" ) @@ -278,7 +278,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"free-model-alias": {"model": "free-model", "hidden": False}}, ) - assert _is_model_cost_zero(model="free-model-alias", llm_router=router) is True + assert is_model_cost_zero(model="free-model-alias", llm_router=router) is True def test_model_group_alias_to_paid_model_enforces_budget(self): """An alias does not turn a priced group into a free one.""" @@ -293,7 +293,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"paid-model-alias": "paid-model"}, ) - assert _is_model_cost_zero(model="paid-model-alias", llm_router=router) is False + assert is_model_cost_zero(model="paid-model-alias", llm_router=router) is False def test_model_group_alias_to_ptu_flat_cost_enforces_budget(self): """A PTU group keeps budget enforced through an alias. @@ -323,8 +323,8 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"ptu-model-alias": "ptu-model"}, ) - assert _is_model_cost_zero(model="ptu-model", llm_router=router) is False - assert _is_model_cost_zero(model="ptu-model-alias", llm_router=router) is False, ( + assert is_model_cost_zero(model="ptu-model", llm_router=router) is False + assert is_model_cost_zero(model="ptu-model-alias", llm_router=router) is False, ( "An aliased PTU group must not be read as free" ) @@ -350,7 +350,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"hidden-alias": {"model": "free-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is True + assert is_model_cost_zero(model="hidden-alias", llm_router=router) is True def test_hidden_model_group_alias_to_paid_model_enforces_budget(self): """A hidden alias to a priced group keeps budget enforced.""" @@ -365,7 +365,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"hidden-paid-alias": {"model": "paid-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="hidden-paid-alias", llm_router=router) is False + assert is_model_cost_zero(model="hidden-paid-alias", llm_router=router) is False def test_dangling_model_group_alias_enforces_budget(self): """An alias pointing at a group that does not exist keeps budget enforced.""" @@ -385,7 +385,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"dangling-alias": "model-that-does-not-exist"}, ) - assert _is_model_cost_zero(model="dangling-alias", llm_router=router) is False + assert is_model_cost_zero(model="dangling-alias", llm_router=router) is False def test_repointed_hidden_alias_does_not_reuse_cached_free_result(self): """Repointing a hidden alias from a free group to a paid group re-evaluates the cost. @@ -415,9 +415,9 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"hidden-alias": {"model": "free-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is True + assert is_model_cost_zero(model="hidden-alias", llm_router=router) is True router.update_settings(model_group_alias={"hidden-alias": {"model": "paid-model", "hidden": True}}) - assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is False + assert is_model_cost_zero(model="hidden-alias", llm_router=router) is False @pytest.mark.parametrize("alias_name_first", [True, False]) def test_alias_shadowing_a_real_group_answers_for_its_target_in_either_order(self, alias_name_first: bool): @@ -456,10 +456,10 @@ class TestUnmappedModelBudgetEnforcement: order = ("ptu-model", "free-model") if alias_name_first else ("free-model", "ptu-model") expected = {"ptu-model": True, "free-model": True} - assert [_is_model_cost_zero(model=name, llm_router=router) for name in order] == [ + assert [is_model_cost_zero(model=name, llm_router=router) for name in order] == [ expected[name] for name in order ] - assert [_is_model_cost_zero(model=name, llm_router=router) for name in order] == [ + assert [is_model_cost_zero(model=name, llm_router=router) for name in order] == [ expected[name] for name in order ], "the cached verdicts must match the first evaluation" @@ -508,8 +508,8 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"free-model": alias_entry}, ) - assert _is_model_cost_zero(model="unpriced-target", llm_router=router) is False - assert _is_model_cost_zero(model="free-model", llm_router=router) is False, ( + assert is_model_cost_zero(model="unpriced-target", llm_router=router) is False + assert is_model_cost_zero(model="free-model", llm_router=router) is False, ( "the alias routes to the unpriced target, so it must be refused like the target by name" ) @@ -546,8 +546,8 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"ptu-model": {"model": "free-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="free-model", llm_router=router) is True - assert _is_model_cost_zero(model="ptu-model", llm_router=router) is True, ( + assert is_model_cost_zero(model="free-model", llm_router=router) is True + assert is_model_cost_zero(model="ptu-model", llm_router=router) is True, ( "the alias routes to the free target, so it must bypass budget like the target by name" ) @@ -579,7 +579,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"free-model": {"model": "paid-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="free-model", llm_router=router) is False + assert is_model_cost_zero(model="free-model", llm_router=router) is False def test_alias_chain_through_a_priced_group_enforces_budget(self): """An alias to a group that is itself an alias key resolves one hop, like the router does. @@ -613,7 +613,7 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"chain-smart": "chain-legacy", "chain-legacy": "free-model"}, ) - assert _is_model_cost_zero(model="chain-smart", llm_router=router) is False + assert is_model_cost_zero(model="chain-smart", llm_router=router) is False @pytest.mark.parametrize("hidden", [False, True], ids=["plain_alias", "hidden_alias"]) def test_alias_to_an_unpriced_group_that_is_also_an_alias_enforces_budget(self, hidden: bool): @@ -634,8 +634,8 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "chain-entry") == UNPRICED_ZERO_COST_MODEL - assert _is_model_cost_zero(model="chain-entry", llm_router=router) is False - assert _is_model_cost_zero(model="chain-middle", llm_router=router) is True, ( + assert is_model_cost_zero(model="chain-entry", llm_router=router) is False + assert is_model_cost_zero(model="chain-middle", llm_router=router) is True, ( "asked by its own name, chain-middle routes to the explicitly free group" ) @@ -655,8 +655,8 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "chain-entry") == "gpt-3.5-turbo" - assert _is_model_cost_zero(model="chain-entry", llm_router=router) is True - assert _is_model_cost_zero(model="chain-middle", llm_router=router) is False, ( + assert is_model_cost_zero(model="chain-entry", llm_router=router) is True + assert is_model_cost_zero(model="chain-middle", llm_router=router) is False, ( "asked by its own name, chain-middle routes to the PTU group" ) @@ -676,7 +676,7 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "chain-entry") == "openai/gpt-4o-mini" - assert _is_model_cost_zero(model="chain-entry", llm_router=router) is False + assert is_model_cost_zero(model="chain-entry", llm_router=router) is False @pytest.mark.parametrize("alias_name", ["openai/smart", "smart"], ids=["alias_on_pattern", "alias_off_pattern"]) def test_alias_chain_served_by_an_explicitly_priced_wildcard_route_enforces_budget(self, alias_name: str): @@ -704,7 +704,7 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, alias_name) == "openai/gpt-4o-mini" - assert _is_model_cost_zero(model=alias_name, llm_router=router) is False + assert is_model_cost_zero(model=alias_name, llm_router=router) is False def test_alias_shadowing_a_free_group_is_judged_by_its_unpriced_target_through_an_alias_chain(self): """A shadowing alias stays enforced when its unpriced target is itself an alias key to a free group.""" @@ -718,7 +718,7 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "shadowed-free") == UNPRICED_ZERO_COST_MODEL - assert _is_model_cost_zero(model="shadowed-free", llm_router=router) is False + assert is_model_cost_zero(model="shadowed-free", llm_router=router) is False @pytest.mark.parametrize("alias_name", ["ollama/fast", "fast"], ids=["alias_on_pattern", "alias_off_pattern"]) def test_alias_to_a_name_served_by_an_explicitly_free_wildcard_route_bypasses_budget(self, alias_name: str): @@ -729,8 +729,8 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, alias_name) == "ollama/llama3" - assert _is_model_cost_zero(model="ollama/llama3", llm_router=router) is True - assert _is_model_cost_zero(model=alias_name, llm_router=router) is True + assert is_model_cost_zero(model="ollama/llama3", llm_router=router) is True + assert is_model_cost_zero(model=alias_name, llm_router=router) is True def test_alias_chain_to_a_name_served_by_an_explicitly_free_wildcard_route_bypasses_budget(self): """An alias to an alias key no deployment is named after reads the wildcard route serving it. @@ -745,8 +745,8 @@ class TestUnmappedModelBudgetEnforcement: ) assert _served_model(router, "ollama/fast") == "ollama/llama3" - assert _is_model_cost_zero(model="ollama/fast", llm_router=router) is True - assert _is_model_cost_zero(model="ollama/llama3", llm_router=router) is False, ( + assert is_model_cost_zero(model="ollama/fast", llm_router=router) is True + assert is_model_cost_zero(model="ollama/llama3", llm_router=router) is False, ( "asked by its own name, ollama/llama3 routes to the unpriced group" ) @@ -769,5 +769,5 @@ class TestUnmappedModelBudgetEnforcement: # Strip the attribute so the helper falls back to the no-cache path. del mock_router._zero_cost_cache - result = _is_model_cost_zero(model="paid-model", llm_router=mock_router) + result = is_model_cost_zero(model="paid-model", llm_router=mock_router) assert result is False diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index e84ecea2acd..a58599d28e9 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -572,7 +572,7 @@ def test_allowed_route_inside_route(user_role, auth_user_id, requested_user_id, def test_read_request_body(): - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body from fastapi import Request payload = "()" * 1000000 @@ -582,7 +582,7 @@ def test_read_request_body(): return payload request.body = return_body - result = _read_request_body(request) + result = read_request_body(request) assert result is not None @@ -814,9 +814,9 @@ def test_is_allowed_route(): ], ) def test_is_user_proxy_admin(user_obj, expected_result): - from litellm.proxy.auth.auth_checks import _is_user_proxy_admin + from litellm.proxy.auth.auth_checks import is_user_proxy_admin - assert _is_user_proxy_admin(user_obj) == expected_result + assert is_user_proxy_admin(user_obj) == expected_result @pytest.mark.parametrize( @@ -846,9 +846,9 @@ def test_is_user_proxy_admin(user_obj, expected_result): ], ) def test_get_user_role(user_obj, expected_role): - from litellm.proxy.auth.user_api_key_auth import _get_user_role + from litellm.proxy.auth.auth_checks import get_user_role - assert _get_user_role(user_obj) == expected_role + assert get_user_role(user_obj) == expected_role @pytest.mark.asyncio diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index 5de1d9300fe..4ba2c2c56f2 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -45,7 +45,7 @@ from litellm.proxy.auth.auth_checks import ( TeamNotFoundError, UserNotFoundError, get_key_object, - _cache_key_object, + cache_key_object, jwt_key_mapping_cache_key, ) from litellm.proxy.auth.route_checks import RouteChecks @@ -59,9 +59,9 @@ from litellm.proxy.auth.user_api_key_auth import ( _reserve_budget_after_common_checks, _route_requires_auth_despite_public, _routing_selector_matches_claim, - _run_centralized_common_checks, + run_centralized_common_checks, _run_post_custom_auth_checks, - _user_api_key_auth_builder, + user_api_key_auth_builder, get_api_key, user_api_key_auth, user_api_key_auth_websocket_for_model, @@ -311,7 +311,7 @@ async def test_should_not_reuse_cached_key_object_for_request_state(): }, ) - await _cache_key_object( + await cache_key_object( hashed_token="cached-token", user_api_key_obj=cached_key, user_api_key_cache=key_cache, @@ -528,7 +528,7 @@ async def test_user_custom_auth_skips_post_custom_auth_checks_by_default(): import litellm import litellm.proxy.proxy_server as _proxy_server_mod from litellm.proxy._types import LitellmUserRoles - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder trusted_token = UserAPIKeyAuth( api_key="sk-custom-auth-trusted", @@ -553,7 +553,7 @@ async def test_user_custom_auth_skips_post_custom_auth_checks_by_default(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key="Bearer sk-custom-auth-trusted", azure_api_key_header="", @@ -586,7 +586,7 @@ async def test_user_custom_auth_runs_post_custom_auth_checks_when_opt_in(): import litellm import litellm.proxy.proxy_server as _proxy_server_mod from litellm.proxy._types import LitellmUserRoles - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder trusted_token = UserAPIKeyAuth( api_key="sk-custom-auth-trusted", @@ -612,7 +612,7 @@ async def test_user_custom_auth_runs_post_custom_auth_checks_when_opt_in(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key="Bearer sk-custom-auth-trusted", azure_api_key_header="", @@ -643,7 +643,7 @@ async def test_enterprise_custom_auth_skips_post_custom_auth_checks_by_default() import litellm import litellm.proxy.proxy_server as _proxy_server_mod from litellm.proxy._types import LitellmUserRoles - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder trusted_token = UserAPIKeyAuth( api_key="sk-enterprise-custom-auth-trusted", @@ -674,7 +674,7 @@ async def test_enterprise_custom_auth_skips_post_custom_auth_checks_by_default() request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key="Bearer sk-enterprise-custom-auth-trusted", azure_api_key_header="", @@ -706,7 +706,7 @@ async def test_enterprise_custom_auth_runs_post_custom_auth_checks_when_opt_in() import litellm import litellm.proxy.proxy_server as _proxy_server_mod from litellm.proxy._types import LitellmUserRoles - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder trusted_token = UserAPIKeyAuth( api_key="sk-enterprise-custom-auth-trusted", @@ -738,7 +738,7 @@ async def test_enterprise_custom_auth_runs_post_custom_auth_checks_when_opt_in() request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key="Bearer sk-enterprise-custom-auth-trusted", azure_api_key_header="", @@ -1158,7 +1158,7 @@ async def test_proxy_admin_expired_key_from_cache(): ProxyException, UserAPIKeyAuth, ) - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token # Create an expired PROXY_ADMIN key @@ -1196,7 +1196,7 @@ async def test_proxy_admin_expired_key_from_cache(): new_callable=AsyncMock, ) as mock_get_key_object, patch( - "litellm.proxy.auth.user_api_key_auth._delete_cache_key_object", + "litellm.proxy.auth.user_api_key_auth.delete_cache_key_object", new_callable=AsyncMock, ) as mock_delete_cache, ): @@ -1232,7 +1232,7 @@ async def test_proxy_admin_expired_key_from_cache(): # Call the auth builder - should raise ProxyException for expired key # Note: api_key needs "Bearer " prefix for get_api_key() to process it correctly with pytest.raises(ProxyException) as exc_info: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", # Add Bearer prefix azure_api_key_header="", @@ -1281,7 +1281,7 @@ async def test_scim_deactivated_user_key_is_rejected(): from fastapi import Request from starlette.datastructures import URL - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-scim-deactivated-user-key" @@ -1346,7 +1346,7 @@ async def test_scim_deactivated_user_key_is_rejected(): ), ): with pytest.raises(ProxyException) as exc_info: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -1371,7 +1371,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker(): from fastapi import Request from starlette.datastructures import URL - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-cached-admin-marker-test" @@ -1424,7 +1424,7 @@ async def test_cached_proxy_admin_key_sets_via_virtual_key_marker(): new_callable=AsyncMock, return_value=cached_token, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -1452,7 +1452,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker(): from starlette.datastructures import URL from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder master_key = "sk-master-key" @@ -1490,7 +1490,7 @@ async def test_master_key_auth_sets_via_virtual_key_marker(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {master_key}", azure_api_key_header="", @@ -1516,7 +1516,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker(): from fastapi import Request from starlette.datastructures import URL - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-via-virtual-key-marker-test" @@ -1576,7 +1576,7 @@ async def test_db_virtual_key_auth_sets_via_virtual_key_marker(): return_value=None, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -1601,7 +1601,7 @@ async def test_auth_prefetches_referenced_objects_only_after_the_key_may_call_th from fastapi import Request from starlette.datastructures import URL - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-prefetch-order-test" @@ -1649,7 +1649,7 @@ async def test_auth_prefetches_referenced_objects_only_after_the_key_may_call_th return_value=valid_token, ), patch( # test-quality-ok: the observable is whether the prefetch runs before or after this check - "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", + "litellm.proxy.auth.user_api_key_auth.enforce_key_and_fallback_model_access", new_callable=AsyncMock, side_effect=None if model_allowed else denied, ), @@ -1660,7 +1660,7 @@ async def test_auth_prefetches_referenced_objects_only_after_the_key_may_call_th "litellm.proxy.auth.user_api_key_auth.get_user_object", new_callable=AsyncMock, return_value=None ), ): - call = _user_api_key_auth_builder( + call = user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -1868,7 +1868,7 @@ async def test_standard_jwt_auth_propagates_user_email(): return_value=mock_jwt_result, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -1938,7 +1938,7 @@ async def test_jwt_auth_propagates_agent_id_to_user_api_key_auth(is_proxy_admin: return_value=mock_jwt_result, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -2452,7 +2452,7 @@ async def test_auto_register_map_existing_key_first_request_runs_key_checks( return_value=reused_key, ), ): - call = _user_api_key_auth_builder( + call = user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -2562,7 +2562,7 @@ async def test_auto_register_first_request_propagates_user_email(active: bool) - ): if not active: with pytest.raises(ProxyException, match="deactivated via SCIM") as exc: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -2574,7 +2574,7 @@ async def test_auto_register_first_request_propagates_user_email(active: bool) - assert int(exc.value.code) == 401 auto_register.assert_not_awaited() return - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -2791,7 +2791,7 @@ async def test_jwt_auto_register_forwards_bound_agent_id(): auto_register, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -3107,7 +3107,7 @@ class TestJWTOAuth2Coexistence: return_value=auto_registered_key, ) as mock_auto_register, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -3187,7 +3187,7 @@ class TestJWTOAuth2Coexistence: return_value=backfilled_user, ) as mock_get_user_object, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -3262,7 +3262,7 @@ class TestJWTOAuth2Coexistence: return_value=other_owner, ) as mock_get_user_object, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -3333,7 +3333,7 @@ class TestJWTOAuth2Coexistence: side_effect=Exception("can't reach database server"), ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -4050,7 +4050,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls(): from starlette.requests import Request from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder _blocking_methods = [ "set_cache", @@ -4142,7 +4142,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls(): return_value=None, ) ) - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4184,7 +4184,7 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): LitellmUserRoles, UserAPIKeyAuth, ) - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder api_key = "sk-test-team-metadata-refresh" @@ -4252,7 +4252,7 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): return_value=fresh_team_obj, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4297,7 +4297,7 @@ async def test_auth_flow_never_persists_fallback_team_object_lit_4391(): from fastapi import HTTPException from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder api_key = "sk-test-lit-4391-no-team-writeback" valid_token = UserAPIKeyAuth( @@ -4358,7 +4358,7 @@ async def test_auth_flow_never_persists_fallback_team_object_lit_4391(): ), ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4400,7 +4400,7 @@ async def test_auth_flow_fallback_team_resolves_object_permission_by_id(): LitellmUserRoles, UserAPIKeyAuth, ) - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder api_key = "sk-test-fallback-team-object-permission" valid_token = UserAPIKeyAuth( @@ -4471,7 +4471,7 @@ async def test_auth_flow_fallback_team_resolves_object_permission_by_id(): return_value=restricted_object_permission, ) as mock_get_object_permission, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4501,7 +4501,7 @@ async def test_auth_flow_fallback_team_object_permission_none_when_unreadable(): from fastapi import HTTPException from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder api_key = "sk-test-fallback-team-object-permission-unreadable" valid_token = UserAPIKeyAuth( @@ -4566,7 +4566,7 @@ async def test_auth_flow_fallback_team_object_permission_none_when_unreadable(): return_value=None, ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -4631,7 +4631,7 @@ async def test_centralized_common_checks_runs_for_standard_auth(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -4690,7 +4690,7 @@ async def test_centralized_common_checks_routes_header_tags_to_litellm_metadata( new_callable=AsyncMock, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data=request_data, @@ -4744,7 +4744,7 @@ async def test_centralized_common_checks_carries_team_and_user_budget_state_on_t "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-5.4-mini"}, @@ -4817,7 +4817,7 @@ async def _run_centralized_checks_with_key_end_user_budget( new_callable=AsyncMock, ) as mock_reserve, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-5.4-mini", "user": request_user or token.end_user_id}, @@ -4998,7 +4998,7 @@ async def test_centralized_common_checks_enforces_team_model_max_budget_from_the new_callable=AsyncMock, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5034,7 +5034,7 @@ async def test_centralized_common_checks_skipped_for_custom_auth_without_flag(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5131,7 +5131,7 @@ async def test_centralized_checks_enforce_token_end_user_budget_against_row_spen ), pytest.raises(litellm.BudgetExceededError) as exc_info, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=_chat_request(), request_data={ @@ -5167,7 +5167,7 @@ async def test_centralized_checks_skip_end_user_lookup_without_a_token_budget(): return_value=0.6, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=_chat_request(), request_data={ @@ -5201,7 +5201,7 @@ async def test_centralized_common_checks_runs_for_custom_auth_with_flag(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5242,7 +5242,7 @@ async def test_centralized_common_checks_runs_for_oauth2_fallback_token(): ), ): with pytest.raises(ProxyException) as exc: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4"}, @@ -5292,7 +5292,7 @@ async def test_centralized_common_checks_tolerates_db_errors_when_fetching_conte new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5350,7 +5350,7 @@ async def test_team_key_vector_store_access_when_team_cannot_be_resolved( "litellm.proxy.auth.user_api_key_auth.get_team_object", AsyncMock(side_effect=team_lookup_error) ) - checks: Final = _run_centralized_common_checks( + checks: Final = run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={ @@ -5434,7 +5434,7 @@ async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_te ), ) - checks: Final = _run_centralized_common_checks( + checks: Final = run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={ @@ -5501,7 +5501,7 @@ async def test_keyless_proxy_admin_keeps_personal_vector_store_grants_under_deny ), ) - checks: Final = _run_centralized_common_checks( + checks: Final = run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={ @@ -5561,7 +5561,7 @@ async def test_centralized_common_checks_propagates_end_user_budget_error(): ) as mock_checks, ): with pytest.raises(litellm.BudgetExceededError): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"user": "alice", "model": "gpt-4o"}, @@ -5624,7 +5624,7 @@ async def test_centralized_common_checks_reserves_request_end_user_budget(): ): assert token.end_user_id is None - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data=request_data, @@ -5674,7 +5674,7 @@ async def test_centralized_common_checks_short_circuits_when_master_key_unset(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -5710,7 +5710,7 @@ async def test_centralized_common_checks_skips_public_routes(): "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -5756,7 +5756,7 @@ async def test_centralized_common_checks_skips_passthrough_endpoint_with_auth_fa "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -5800,7 +5800,7 @@ async def test_centralized_common_checks_runs_for_passthrough_endpoint_with_auth "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock, ) as mock_checks: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -5859,7 +5859,7 @@ async def test_centralized_common_checks_master_key_admin_overrides_db_user_role new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"team_id": "t1", "max_budget": 10}, @@ -5910,7 +5910,7 @@ async def test_centralized_common_checks_http_exception_without_team_id(): ): # Should NOT raise AssertionError from _team_obj_from_token; # should proceed with team_object=None. - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -5999,7 +5999,7 @@ async def test_centralized_common_checks_team_404_does_not_zero_other_contexts() new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"user": "alice", "model": "gpt-4o"}, @@ -6057,7 +6057,7 @@ async def test_centralized_common_checks_unresolvable_team_without_grant_is_refu side_effect=team_read_failure, ): with pytest.raises(HTTPException) as exc_info: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4.1"}, @@ -6111,7 +6111,7 @@ async def test_centralized_common_checks_absent_team_refused_despite_db_unavaila side_effect=team_absent, ): with pytest.raises(HTTPException) as exc_info: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4.1"}, @@ -6164,7 +6164,7 @@ async def test_centralized_common_checks_unreadable_team_keeps_db_unavailable_op _capturing_common_checks, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4.1"}, @@ -6215,7 +6215,7 @@ async def test_centralized_common_checks_unresolvable_team_with_grant_enforces_i side_effect=HTTPException(status_code=404, detail={"error": "team unreadable"}), ): if is_granted: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": requested_model}, @@ -6223,7 +6223,7 @@ async def test_centralized_common_checks_unresolvable_team_with_grant_enforces_i ) else: with pytest.raises(ProxyException) as exc_info: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": requested_model}, @@ -6283,7 +6283,7 @@ async def test_centralized_common_checks_ui_sentinel_team_vouches_despite_absent _capturing_common_checks, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -6344,7 +6344,7 @@ async def test_centralized_common_checks_ui_sentinel_team_skips_db_lookup(): _capturing_common_checks, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={}, @@ -6371,7 +6371,7 @@ async def test_builder_ui_sentinel_team_never_hits_get_team_object(): # test-qu from starlette.datastructures import URL from litellm.proxy._types import UI_TEAM_ID - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.proxy_server import hash_token api_key = "sk-test-ui-session-key" @@ -6419,7 +6419,7 @@ async def test_builder_ui_sentinel_team_never_hits_get_team_object(): # test-qu new_callable=AsyncMock, ) as mock_get_team_object, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -6503,7 +6503,7 @@ async def test_centralized_common_checks_user_http_exception_isolates_to_user_on new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"user": "alice", "model": "gpt-4o"}, @@ -6567,7 +6567,7 @@ async def test_centralized_common_checks_backfills_org_id_from_team(key_org_id, side_effect=lambda **kw: org_id_seen_by_common_checks.append(kw["valid_token"].org_id), ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -6713,14 +6713,14 @@ async def test_centralized_common_checks_inherits_org_identity( if expect_lookup_error: with pytest.raises(ConnectionRefusedError, match="db unavailable"): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, route="/chat/completions", ) else: - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -6809,7 +6809,7 @@ async def test_cli_session_token_org_backfilled_from_team(monkeypatch): new_callable=AsyncMock, ), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -6851,7 +6851,7 @@ async def test_centralized_common_checks_org_backfill_survives_team_fetch_failur new_callable=AsyncMock, ) as mock_checks, ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-4o"}, @@ -6878,7 +6878,7 @@ async def test_master_key_auth_substitutes_alias_for_api_key(): from starlette.datastructures import URL from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.utils import hash_token import litellm.proxy.proxy_server as _proxy_server_mod @@ -6893,7 +6893,7 @@ async def test_master_key_auth_substitutes_alias_for_api_key(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {master_key}", azure_api_key_header="", @@ -6951,12 +6951,12 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it(): # auth state machine; we only care about the wrapper's safety net. with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, return_value=builder_token, ), patch( - "litellm.proxy.auth.user_api_key_auth._run_centralized_common_checks", + "litellm.proxy.auth.user_api_key_auth.run_centralized_common_checks", new_callable=AsyncMock, ), patch( @@ -7004,12 +7004,12 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder() setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, return_value=builder_token, ), patch( - "litellm.proxy.auth.user_api_key_auth._run_centralized_common_checks", + "litellm.proxy.auth.user_api_key_auth.run_centralized_common_checks", new_callable=AsyncMock, ), patch( @@ -7060,12 +7060,12 @@ async def test_user_api_key_auth_authenticates_before_raising_malformed_body_err setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, return_value=builder_token, ) as mock_builder, patch( - "litellm.proxy.auth.user_api_key_auth._run_centralized_common_checks", + "litellm.proxy.auth.user_api_key_auth.run_centralized_common_checks", new_callable=AsyncMock, ) as mock_common_checks, patch( @@ -7121,7 +7121,7 @@ async def _run_auth_with_malformed_body(post_call_failure_hook): setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, return_value=builder_token, ), @@ -7193,7 +7193,7 @@ async def test_user_api_key_auth_malformed_body_with_rejected_key_still_returns_ setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, side_effect=ProxyException( message="Authentication Error, invalid key", @@ -7246,7 +7246,7 @@ async def test_user_api_key_auth_does_not_double_log_a_malformed_body_from_a_rej setattr(_proxy_server_mod, k, v) with ( patch( - "litellm.proxy.auth.user_api_key_auth._user_api_key_auth_builder", + "litellm.proxy.auth.user_api_key_auth.user_api_key_auth_builder", new_callable=AsyncMock, side_effect=ProxyException( message="Authentication Error, invalid key", @@ -7298,7 +7298,7 @@ async def _run_builder_with_key_lookup(get_key_object_mock): from starlette.datastructures import URL import litellm.proxy.proxy_server as _proxy_server_mod - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder attrs = _proxy_attrs_for_db_lookup() originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} @@ -7316,7 +7316,7 @@ async def _run_builder_with_key_lookup(get_key_object_mock): "litellm.proxy.auth.auth_exception_handler.seed_request_identity", ), ): - return await _user_api_key_auth_builder( + return await user_api_key_auth_builder( request=request, api_key="Bearer sk-db-lookup-test", azure_api_key_header="", @@ -7628,7 +7628,7 @@ async def test_non_admin_cli_session_token_reaches_production_auth_path(monkeypa side_effect=__import__("fastapi").HTTPException(status_code=404), ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {cli_token}", azure_api_key_header="", @@ -7715,7 +7715,7 @@ async def _authenticate_session_token_against_db( AsyncMock(return_value=membership_row), ), ): - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=request, api_key=f"Bearer {cli_token}", azure_api_key_header="", @@ -7925,7 +7925,7 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): from starlette.datastructures import URL import litellm.proxy.proxy_server as _proxy_server_mod - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_builder from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.proxy_server import hash_token @@ -7969,10 +7969,10 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") with patch( - "litellm.proxy.auth.resolvers.store._fetch_key_object_from_db_with_reconnect", + "litellm.proxy.auth.resolvers.store.fetch_key_object_from_db_with_reconnect", fetch_from_db, ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8388,7 +8388,7 @@ async def test_global_proxy_spend_reads_resettable_proxy_budget_row(): loaded from the MonthlyGlobalSpend view, whose window is hardcoded to a trailing 30 days and never resets on the configured duration.""" from litellm.proxy.auth.user_api_key_auth import ( - _fetch_global_spend_with_event_coordination, + fetch_global_spend_with_event_coordination, ) from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -8400,7 +8400,7 @@ async def test_global_proxy_spend_reads_resettable_proxy_budget_row(): side_effect=AssertionError("global spend must not be loaded from the fixed-30d MonthlyGlobalSpend view") ) - result = await _fetch_global_spend_with_event_coordination( + result = await fetch_global_spend_with_event_coordination( cache_key="default_user_id:spend", user_api_key_cache=UserApiKeyCache(), prisma_client=prisma_client, @@ -8415,14 +8415,14 @@ async def test_global_proxy_spend_none_when_proxy_budget_row_missing(): """Before the startup upsert creates the aggregate row, enforcement must see None (no cap applied) rather than raising.""" from litellm.proxy.auth.user_api_key_auth import ( - _fetch_global_spend_with_event_coordination, + fetch_global_spend_with_event_coordination, ) from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache prisma_client = MagicMock() prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) - result = await _fetch_global_spend_with_event_coordination( + result = await fetch_global_spend_with_event_coordination( cache_key="default_user_id:spend", user_api_key_cache=UserApiKeyCache(), prisma_client=prisma_client, @@ -8463,7 +8463,7 @@ async def test_temp_budget_increase_applied_for_cached_key(): ) user_api_key_cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=cached_key, user_api_key_cache=user_api_key_cache, @@ -8487,13 +8487,13 @@ async def test_temp_budget_increase_applied_for_cached_key(): patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj), patch( - "litellm.proxy.auth.user_api_key_auth._virtual_key_max_budget_alert_check", + "litellm.proxy.auth.user_api_key_auth.virtual_key_max_budget_alert_check", new_callable=AsyncMock, ), ): results = tuple( [ - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8535,7 +8535,7 @@ async def test_cached_key_team_member_budget_blocks_at_exact_cap(team_member_spe max_budget = 2.4 user_api_key_cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=UserAPIKeyAuth( token=hashed_token, @@ -8577,7 +8577,7 @@ async def test_cached_key_team_member_budget_blocks_at_exact_cap(team_member_spe proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) async def _auth(): - return await _user_api_key_auth_builder( + return await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8637,7 +8637,7 @@ async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset team_member_spend = 2.5 user_api_key_cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=UserAPIKeyAuth( token=hashed_token, @@ -8683,7 +8683,7 @@ async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) async def _auth(): - return await _user_api_key_auth_builder( + return await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8726,7 +8726,7 @@ async def _authenticate_and_authorize(mock_request, api_key): from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request request_data = {"model": "claude-sonnet-5", "messages": [{"role": "user", "content": "hi"}]} - auth_obj = await _user_api_key_auth_builder( + auth_obj = await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -8773,7 +8773,7 @@ async def test_cached_key_team_member_budget_emails_configured_thresholds( alert_emails = {"50": [], "100": ["finance@example.com"]} user_api_key_cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hashed_token, user_api_key_obj=UserAPIKeyAuth( token=hashed_token, @@ -8900,7 +8900,7 @@ async def _proxy_exception_for_key( patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), ): with pytest.raises(ProxyException) as exc_info: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=mock_request, api_key=f"Bearer {api_key}", azure_api_key_header="", @@ -9208,7 +9208,7 @@ async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_a new_callable=AsyncMock, return_value=builder_result, ): - token = await _user_api_key_auth_builder( + token = await user_api_key_auth_builder( request=request, api_key="Bearer header.payload.signature", azure_api_key_header="", @@ -9243,7 +9243,7 @@ async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_a ) async def test_claude_view_normalizes_before_model_access(monkeypatch, route): from starlette.requests import Request - from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access + from litellm.proxy.auth.user_api_key_auth import enforce_key_and_fallback_model_access source = "foo[1m]" encoded = "claude-router-" + source.encode().hex() + "[1m]" @@ -9254,7 +9254,7 @@ async def test_claude_view_normalizes_before_model_access(monkeypatch, route): data = {"model": encoded, "messages": [{"role": "user", "content": "hi"}]} request = Request({"type": "http", "method": "POST", "path": route, "headers": [], "query_string": b""}) token = UserAPIKeyAuth(models=[source]) - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=token, request_data=data, route=route, @@ -9267,7 +9267,7 @@ async def test_claude_view_normalizes_before_model_access(monkeypatch, route): assert json.loads(await request.body())["model"] == source assert request.scope["parsed_body"][1]["model"] == source with pytest.raises(ProxyException): - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=UserAPIKeyAuth(models=["other"]), request_data=data, route=route, @@ -9668,7 +9668,7 @@ async def test_auth_flow_enters_virtual_key_mapping_when_only_an_issuer_configur side_effect=AssertionError("standard JWT auth must not run for a mapped virtual key"), ), ): - result = await _user_api_key_auth_builder( + result = await user_api_key_auth_builder( request=mock_request, api_key=jwt_token, azure_api_key_header="", @@ -9712,9 +9712,9 @@ def _alias_request(route: str, data: dict, content_type: str = "application/json async def _enforce_alias_access(token: UserAPIKeyAuth, data: dict, route: str, request, router: litellm.Router): - from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access + from litellm.proxy.auth.user_api_key_auth import enforce_key_and_fallback_model_access - await _enforce_key_and_fallback_model_access( + await enforce_key_and_fallback_model_access( valid_token=token, request_data=data, route=route, @@ -9783,18 +9783,18 @@ async def test_router_settings_model_group_alias_leaves_form_bodies_alone(monkey @pytest.mark.asyncio async def test_router_settings_model_group_alias_rewrite_keeps_query_params_out_of_body(monkeypatch): """LIT-3054: auth merges query params into its own copy of the body; the rewrite must not forward them.""" - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body, populate_request_with_path_params + from litellm.proxy.common_utils.http_parsing_utils import read_request_body, populate_request_with_path_params router = _alias_router() monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) body = {"model": "AgentX-LLM", "messages": [{"role": "user", "content": "hi"}]} request = _alias_request("/v1/chat/completions", body) request.scope["query_string"] = b"api-version=2024-10-21&stream=true" - data = populate_request_with_path_params(request_data=await _read_request_body(request), request=request) + data = populate_request_with_path_params(request_data=await read_request_body(request), request=request) assert data["api-version"] == "2024-10-21" token = _alias_token(monkeypatch, "key", {"AgentX-LLM": "claude-haiku"}, ["claude-haiku"]) await _enforce_alias_access(token, data, "/v1/chat/completions", request, router) - downstream = await _read_request_body(request) + downstream = await read_request_body(request) assert downstream == {**body, "model": "claude-haiku"} assert json.loads(await request.body()) == downstream assert await request.json() == downstream @@ -10003,7 +10003,7 @@ async def test_admission_and_budget_reservation_read_the_key_spend_counter_with_ ), spend_counter_batch_scope(redis), ): - await _run_centralized_common_checks( + await run_centralized_common_checks( user_api_key_auth_obj=token, request=request, request_data={"model": "gpt-5.4-mini", "messages": [{"role": "user", "content": "hi"}]}, @@ -10242,7 +10242,7 @@ async def test_managed_jwt_cannot_be_downgraded_into_virtual_key_mapping(monkeyp monkeypatch.setattr(proxy_server, name, value) for _ in range(2): with pytest.raises(ProxyException) as failure: - await _user_api_key_auth_builder( + await user_api_key_auth_builder( request=_alias_request("/v1/chat/completions", {}), api_key="Bearer verified.jwt.token", azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None, request_data={}, @@ -10271,7 +10271,7 @@ async def test_virtual_key_cannot_enter_checks_as_an_identity_managed_actor(monk client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) monkeypatch.setattr(proxy_server, "prisma_client", client) checks: Final = AsyncMock() - monkeypatch.setattr(auth_module, "_run_centralized_common_checks", checks) + monkeypatch.setattr(auth_module, "run_centralized_common_checks", checks) monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock(return_value=None))) data: Final = {"model": "allowed", "messages": [{"role": "user", "content": "hello"}]} request: Final = _alias_request("/v1/chat/completions", data) @@ -10324,7 +10324,7 @@ async def test_custom_auth_grants_reach_managed_targets_without_a_virtual_key_ro monkeypatch.setattr(module, "enterprise_custom_auth", custom if enterprise else None) monkeypatch.setattr(agent_registry, "global_agent_registry", registry) monkeypatch.setattr(litellm, "enable_post_custom_auth_checks", False, raising=False) - admitted: Final = await _user_api_key_auth_builder( + admitted: Final = await user_api_key_auth_builder( request=_alias_request("/a2a/target/message/send", {}), api_key=f"Bearer {credential}", azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None, request_data={}, @@ -10346,7 +10346,7 @@ async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(mon monkeypatch.setattr(proxy_server, name, value) module: Final = importlib.import_module("litellm.proxy.auth.user_api_key_auth") monkeypatch.setattr(module, "enterprise_custom_auth", custom) - admitted: Final = await _user_api_key_auth_builder( + admitted: Final = await user_api_key_auth_builder( request=_alias_request("/v1/chat/completions", {}), api_key="Bearer external-credential", azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None, request_data={}, diff --git a/tests/unit/proxy/batches_endpoints/test_endpoints.py b/tests/unit/proxy/batches_endpoints/test_endpoints.py index e2ef8e1c789..3469dbee748 100644 --- a/tests/unit/proxy/batches_endpoints/test_endpoints.py +++ b/tests/unit/proxy/batches_endpoints/test_endpoints.py @@ -206,7 +206,7 @@ def harness(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) with ExitStack() as stack: - stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context(patch.object(endpoints, "read_request_body", read_body)) stack.enter_context( patch.object( ProxyBaseLLMRequestProcessing, @@ -593,7 +593,7 @@ async def test_create__unified_file_id_single_model_disables_cross_model_fallbac }, ) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"]), ): resp = await call_create(harness) @@ -620,7 +620,7 @@ async def test_create__unified_file_id_not_exactly_one_model_400(harness, models }, ) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=models), ): with pytest.raises(ProxyException) as exc: @@ -662,7 +662,7 @@ async def test_create__unified_file_id_resolves_real_storage_url(harness): fake_repo_cls = MagicMock(return_value=fake_repo_instance) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]), patch.object(proxy_server, "prisma_client", MagicMock()), patch.object(endpoints, "ManagedFileRepository", fake_repo_cls), @@ -696,7 +696,7 @@ async def test_create__unified_file_id_db_error_falls_back_to_raw_id(harness): fake_repo_cls = MagicMock(return_value=fake_repo_instance) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]), patch.object(proxy_server, "prisma_client", MagicMock()), patch.object(endpoints, "ManagedFileRepository", fake_repo_cls), @@ -728,7 +728,7 @@ async def test_create__multi_model_unified_file_with_loadbalancing_keeps_router_ with ( patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True), - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["model-a", "model-b"]), ): await call_create(harness) @@ -759,7 +759,7 @@ async def test_create__unified_file_id_missing_row_falls_back_to_raw_id(harness) fake_repo_cls = MagicMock(return_value=fake_repo_instance) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]), patch.object(proxy_server, "prisma_client", MagicMock()), patch.object(endpoints, "ManagedFileRepository", fake_repo_cls), @@ -792,7 +792,7 @@ async def test_create__unified_file_id_legacy_row_without_storage_url_dispatches fake_repo_cls = MagicMock(return_value=fake_repo_instance) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]), patch.object(proxy_server, "prisma_client", MagicMock()), patch.object(endpoints, "ManagedFileRepository", fake_repo_cls), @@ -946,7 +946,7 @@ async def test_create__model_encoded_beats_unified(harness): }, ) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["something-else"]), ): await call_create(harness) @@ -1596,7 +1596,7 @@ async def test_retrieve__model_encoded_beats_loadbalancing(retrieve_harness): @pytest.mark.asyncio async def test_retrieve__unified_batch_id_routes_to_router(retrieve_harness): - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): resp = await call_retrieve(retrieve_harness, "batch-unified-blob") # DISPATCH - router fired, direct litellm did not. @@ -1762,7 +1762,7 @@ async def test_retrieve__db_terminal_unified_resolves_file_ids(retrieve_harness) db_batch_object = MagicMock() retrieve_harness.get_batch_from_db.return_value = (db_batch_object, db_response) - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): await call_retrieve(retrieve_harness, "batch-unified-blob") # Terminal short-circuit still registers/normalizes raw provider file ids. @@ -1960,7 +1960,7 @@ def list_harness(): litellm_alist = AsyncMock(return_value=FakeListPage([])) with ExitStack() as stack: - stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context(patch.object(endpoints, "read_request_body", read_body)) stack.enter_context( patch.object( ProxyBaseLLMRequestProcessing, @@ -2495,7 +2495,7 @@ async def test_cancel__model_encoded_id_forwards_deployment_model(cancel_harness @pytest.mark.asyncio async def test_cancel__model_encoded_beats_unified(cancel_harness): - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): await call_cancel(cancel_harness, AZURE_BATCH_ID) assert cancel_harness.litellm_acancel.call_count == 1 @@ -2511,7 +2511,7 @@ async def test_cancel__model_encoded_beats_unified(cancel_harness): @pytest.mark.asyncio async def test_cancel__unified_batch_id_routes_to_router(cancel_harness): - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): resp = await call_cancel(cancel_harness, "batch-unified-blob") # DISPATCH - router fired, litellm did not, no creds lookup. @@ -2537,7 +2537,7 @@ async def test_cancel__db_write_receives_caller_auth(cancel_harness): """update_batch_in_database can only mint managed IDs for a cancelled batch's output files when it has an auth context, so cancel must forward the caller's.""" caller = UserAPIKeyAuth(api_key="sk-test", user_id="user-cancel-1") - with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): + with patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID): await call_cancel(cancel_harness, "batch-unified-blob", user=caller) assert cancel_harness.update_batch_in_db.call_args.kwargs["user_api_key_dict"] is caller @@ -2548,7 +2548,7 @@ async def test_cancel__unified_missing_model_id_400(cancel_harness): # unified id with no model_id segment -> get_model_id returns None -> 400. with patch.object( endpoints, - "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", return_value="litellm_proxy;llm_batch_id:batch-xyz", ): with pytest.raises(ProxyException) as exc: @@ -2563,7 +2563,7 @@ async def test_cancel__unified_missing_model_id_400(cancel_harness): async def test_cancel__unified_no_router_500(cancel_harness): with ( patch.object(proxy_server, "llm_router", None), - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID), ): with pytest.raises(ProxyException) as exc: await call_cancel(cancel_harness, "batch-unified-blob") @@ -2769,7 +2769,7 @@ async def test_create__unified_no_router_500(harness): }, ) with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value="unified-xyz"), patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"]), patch.object(proxy_server, "llm_router", None), ): @@ -2782,7 +2782,7 @@ async def test_create__unified_no_router_500(harness): @pytest.mark.asyncio async def test_retrieve__unified_no_router_500(retrieve_harness): with ( - patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID), + patch.object(endpoints, "is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID), patch.object(proxy_server, "llm_router", None), ): with pytest.raises(ProxyException) as exc: 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 ad99cebc342..00f84e75664 100644 --- a/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/unit/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -30,7 +30,7 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import ( upstream_lacks_files_api, ) from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, get_batch_id_from_unified_batch_id, is_litellm_executed_batch, ) @@ -1129,7 +1129,7 @@ async def test_run_completes_under_the_real_hooks_base64_batch_id() -> None: harness = make_runner(store_factory=RealIdManagedBatchStore) created, finished = await harness.create_and_finish() - assert _is_base64_encoded_unified_file_id(created.id) + assert is_base64_encoded_unified_file_id(created.id) assert finished.status == "completed" assert [call.model_object_id.startswith("litellm_batch_") for call in harness.store.calls] == [True] assert [write.unified_object_id for write in harness.table.writes] == [created.id] * 3 diff --git a/tests/unit/proxy/common_utils/test_cache_aware_routing.py b/tests/unit/proxy/common_utils/test_cache_aware_routing.py index 000d6773d04..24fc1f22a6a 100644 --- a/tests/unit/proxy/common_utils/test_cache_aware_routing.py +++ b/tests/unit/proxy/common_utils/test_cache_aware_routing.py @@ -363,7 +363,7 @@ async def test_router_applies_opt_in_and_preserves_failure_semantics( from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy import proxy_server from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import ProxyLogging from litellm.router_strategy.complexity_router.context_compaction import initialize_compaction_state @@ -390,7 +390,7 @@ async def test_router_applies_opt_in_and_preserves_failure_semantics( ] ) logging: Final = ProxyLogging(UserApiKeyCache()) - logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + logging.proxy_hook_mapping["parallel_request_limiter"] = PROXY_MaxParallelRequestsHandler_v3( logging.internal_usage_cache ) await _observed(logging.internal_usage_cache.dual_cache, expires_at=1e100) 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 d970978704e..87df82c01d0 100644 --- a/tests/unit/proxy/common_utils/test_check_batch_cost.py +++ b/tests/unit/proxy/common_utils/test_check_batch_cost.py @@ -729,7 +729,7 @@ class TestCheckBatchCost: let the original bug ship undetected. """ import litellm - from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() @@ -773,7 +773,7 @@ class TestCheckBatchCost: decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;" - db_logger = _ProxyDBLogger() + db_logger = ProxyDBLogger() mock_update_database = AsyncMock() # Unlike the other tests in this file, this one runs the real @@ -1253,6 +1253,7 @@ class TestCheckBatchCost: mock_response = MagicMock() mock_response.status = terminal_status mock_response.output_file_id = "file-output-123" + mock_response.error_file_id = None mock_response.model_dump_json.return_value = f'{{"id":"batch-1","status":"{terminal_status}"}}' mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) @@ -2272,13 +2273,13 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: @pytest.mark.asyncio async def test_target_model_names_comes_from_input_file_not_provider_model(self): from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, get_models_from_unified_file_id, ) 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) + decoded = is_base64_encoded_unified_file_id(output_file_id) assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP] assert "gpt-5.5" not in decoded @@ -2307,13 +2308,13 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: @pytest.mark.asyncio async def test_falls_back_to_deployment_model_group_without_managed_input_file(self): from litellm.proxy.openai_files_endpoints.common_utils import ( - _is_base64_encoded_unified_file_id, + is_base64_encoded_unified_file_id, get_models_from_unified_file_id, ) output_file_id = await self._run(self._job("file-raw-provider-input")) - decoded = _is_base64_encoded_unified_file_id(output_file_id) + decoded = is_base64_encoded_unified_file_id(output_file_id) assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP] diff --git a/tests/unit/proxy/common_utils/test_encrypt_decrypt_utils.py b/tests/unit/proxy/common_utils/test_encrypt_decrypt_utils.py index 7c07f19d72f..1163090c510 100644 --- a/tests/unit/proxy/common_utils/test_encrypt_decrypt_utils.py +++ b/tests/unit/proxy/common_utils/test_encrypt_decrypt_utils.py @@ -13,7 +13,7 @@ import pytest from litellm.proxy import proxy_server from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _V2_GCM_PREFIX, + V2_GCM_PREFIX, decrypt_bearer_token, decrypt_if_encrypted_with, decrypt_value_helper, @@ -43,7 +43,7 @@ def test_aes_gcm_round_trip(monkeypatch): ct = encrypt_value_helper("super-secret") - assert ct.startswith(_V2_GCM_PREFIX) + assert ct.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(ct, key="t") == "super-secret" @@ -51,7 +51,7 @@ def test_default_is_legacy_algorithm(monkeypatch): """With no config, writes stay on the legacy algorithm (no v2: marker).""" ct = encrypt_value_helper("legacy-secret") - assert not ct.startswith(_V2_GCM_PREFIX) + assert not ct.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(ct, key="t") == "legacy-secret" @@ -62,12 +62,12 @@ def test_legacy_nacl_value_still_decrypts_after_flag_flip(monkeypatch): flipping the flag forward never strands previously-written data. """ legacy = encrypt_value_helper("legacy-secret") # default = xsalsa20 - assert not legacy.startswith(_V2_GCM_PREFIX) + assert not legacy.startswith(V2_GCM_PREFIX) _use_aes(monkeypatch) # New writes are now AES, but the old value must still come back. assert decrypt_value_helper(legacy, key="t") == "legacy-secret" - assert encrypt_value_helper("fresh").startswith(_V2_GCM_PREFIX) + assert encrypt_value_helper("fresh").startswith(V2_GCM_PREFIX) def test_v2_prefix_is_idempotent_marker(monkeypatch): @@ -79,11 +79,11 @@ def test_v2_prefix_is_idempotent_marker(monkeypatch): _use_aes(monkeypatch) ct = encrypt_value_helper("secret") - assert ct.startswith(_V2_GCM_PREFIX) + assert ct.startswith(V2_GCM_PREFIX) # Round-tripping does not change the plaintext, and the marker is stable. again = encrypt_value_helper(decrypt_value_helper(ct, key="t")) - assert again.startswith(_V2_GCM_PREFIX) + assert again.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(again, key="t") == "secret" @@ -91,7 +91,7 @@ def test_aes_decrypt_failure_returns_none_not_raise(monkeypatch): """Decrypt contract preserved: a garbled v2 value returns None, never raises.""" _use_aes(monkeypatch) - garbled = _V2_GCM_PREFIX + "not-valid-base64-or-ciphertext!!!" + garbled = V2_GCM_PREFIX + "not-valid-base64-or-ciphertext!!!" # exception_type="debug" exercises the swallow path; must not raise. assert decrypt_value_helper(garbled, key="t", exception_type="debug") is None @@ -100,7 +100,7 @@ def test_aes_decrypt_failure_returns_original_when_requested(monkeypatch): """With return_original_value=True a bad v2 value comes back as-is, not None.""" _use_aes(monkeypatch) - garbled = _V2_GCM_PREFIX + "###" + garbled = V2_GCM_PREFIX + "###" assert decrypt_value_helper(garbled, key="t", exception_type="debug", return_original_value=True) == garbled @@ -109,7 +109,7 @@ def test_empty_string_round_trips_under_aes(monkeypatch): _use_aes(monkeypatch) ct = encrypt_value_helper("") - assert ct.startswith(_V2_GCM_PREFIX) + assert ct.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(ct, key="t") == "" @@ -121,7 +121,7 @@ def test_callback_prefix_composes_with_v2(monkeypatch): helper is ``v2:gcm:...``. Ordering must work end to end. """ from litellm.proxy.common_utils.callback_utils import ( - _CALLBACK_VAR_ENCRYPTED_PREFIX, + CALLBACK_VAR_ENCRYPTED_PREFIX, _decrypt_or_passthrough, _encrypt_if_plaintext, ) @@ -131,9 +131,9 @@ def test_callback_prefix_composes_with_v2(monkeypatch): # "gcs_path_service_account" is a known-sensitive callback key. stored = _encrypt_if_plaintext("gcs_path_service_account", "my-sa-secret") - assert stored.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX) - inner = stored[len(_CALLBACK_VAR_ENCRYPTED_PREFIX) :] - assert inner.startswith(_V2_GCM_PREFIX) + assert stored.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) + inner = stored[len(CALLBACK_VAR_ENCRYPTED_PREFIX) :] + assert inner.startswith(V2_GCM_PREFIX) assert _decrypt_or_passthrough("gcs_path_service_account", stored) == "my-sa-secret" @@ -142,7 +142,7 @@ def test_unknown_algorithm_falls_back_to_legacy(monkeypatch): monkeypatch.setattr(proxy_server, "general_settings", {"encryption_algorithm": "rot13"}) ct = encrypt_value_helper("secret") - assert not ct.startswith(_V2_GCM_PREFIX) + assert not ct.startswith(V2_GCM_PREFIX) assert decrypt_value_helper(ct, key="t") == "secret" @@ -256,7 +256,7 @@ def test_stored_value_is_not_a_bearer_token_even_when_reshaped(monkeypatch, use_ _use_aes(monkeypatch) stored = encrypt_value_helper("stored-secret") - for candidate in (stored, "kind_a_" + stored.removeprefix(_V2_GCM_PREFIX).rstrip("=")): + for candidate in (stored, "kind_a_" + stored.removeprefix(V2_GCM_PREFIX).rstrip("=")): assert decrypt_bearer_token(candidate, prefix="kind_a_") is None diff --git a/tests/unit/proxy/common_utils/test_http_parsing_utils.py b/tests/unit/proxy/common_utils/test_http_parsing_utils.py index 222a2cb329b..7677837aa0f 100644 --- a/tests/unit/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/unit/proxy/common_utils/test_http_parsing_utils.py @@ -17,11 +17,11 @@ import litellm.proxy.common_utils.http_parsing_utils as http_parsing_utils from litellm.proxy._types import ProxyException from litellm.proxy.common_utils.http_parsing_utils import ( _is_form_content_type, - _read_request_body, - _safe_get_request_headers, + read_request_body, + safe_get_request_headers, _safe_get_request_parsed_body, - _safe_get_request_query_params, - _safe_set_request_parsed_body, + safe_get_request_query_params, + safe_set_request_parsed_body, coerce_numeric_form_fields, get_form_data, get_request_body, @@ -74,8 +74,8 @@ async def test_read_request_body_marks_body_received_once_with_its_size(monkeypa body: Final = orjson.dumps({"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "x" * 4096}]}) request: Final = _starlette_request(body, "application/json") - assert await _read_request_body(request) == orjson.loads(body) - assert await _read_request_body(request) == orjson.loads(body) + assert await read_request_body(request) == orjson.loads(body) + assert await read_request_body(request) == orjson.loads(body) assert events == [("litellm.request.body_received", {"litellm.request.body_bytes": len(body)})] @@ -92,9 +92,9 @@ async def test_read_request_body_marks_body_received_for_binary_and_form_bodies( form: Final = b"model=whisper-1&language=en" form_type: Final = "application/x-www-form-urlencoded" - await _read_request_body(_starlette_request(protobuf, "application/x-protobuf")) - await _read_request_body(_starlette_request(form, form_type, content_length=str(len(form)))) - await _read_request_body(_starlette_request(form, form_type)) + await read_request_body(_starlette_request(protobuf, "application/x-protobuf")) + await read_request_body(_starlette_request(form, form_type, content_length=str(len(form)))) + await read_request_body(_starlette_request(form, form_type)) assert events == [ ("litellm.request.body_received", {"litellm.request.body_bytes": len(protobuf)}), @@ -108,7 +108,7 @@ async def test_read_raw_json_body_returns_the_bytes_the_parsed_body_came_from(): body = b'{"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}]}' request = _starlette_request(body, "application/json") - assert await _read_request_body(request) == orjson.loads(body) + assert await read_request_body(request) == orjson.loads(body) assert await read_raw_json_body(request) == body @@ -124,7 +124,7 @@ async def test_read_raw_json_body_is_none_until_the_body_has_been_parsed(): async def test_read_raw_json_body_is_none_for_form_bodies(): request = _starlette_request(b"model=claude-sonnet-4-5", "application/x-www-form-urlencoded") - assert await _read_request_body(request) == {"model": "claude-sonnet-4-5"} + assert await read_request_body(request) == {"model": "claude-sonnet-4-5"} assert await read_raw_json_body(request) is None @@ -136,7 +136,7 @@ async def test_protobuf_body_is_not_parsed_as_json(content_type): body = b"\n\xa2\x01\n\x1c\n\x0cservice.name\x12\x0c\n\nswarm\xed\xa0\x80\xff" request = _starlette_request(body, content_type) - assert await _read_request_body(request) == {} + assert await read_request_body(request) == {} assert await request.body() == body # body is still readable by the endpoint @@ -144,7 +144,7 @@ async def test_protobuf_body_is_not_parsed_as_json(content_type): async def test_gzipped_json_trace_body_survives_auth_pre_read(): body = gzip.compress(b'{"resourceSpans": []}') request = _starlette_request(body, "application/json", "/v1/traces", "gzip") - assert await _read_request_body(request) == {} + assert await read_request_body(request) == {} assert await request.body() == body @@ -170,7 +170,7 @@ async def test_request_body_caching(): mock_request.scope = {} # First call should parse the body - result1 = await _read_request_body(mock_request) + result1 = await read_request_body(mock_request) assert result1 == test_data assert "parsed_body" in mock_request.scope assert mock_request.scope["parsed_body"] == (("key",), {"key": "value"}) @@ -182,7 +182,7 @@ async def test_request_body_caching(): mock_request.body.reset_mock() # Second call should use the cached body - result2 = await _read_request_body(mock_request) + result2 = await read_request_body(mock_request) assert result2 == {"key": "value"} # Verify the body was not read again @@ -205,7 +205,7 @@ async def test_form_data_parsing(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the form data was correctly parsed assert result == test_data @@ -256,7 +256,7 @@ async def test_form_data_with_json_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the metadata was parsed from JSON string to dict assert "metadata" in result @@ -298,7 +298,7 @@ async def test_form_data_with_invalid_json_metadata(): # Should raise JSONDecodeError when trying to parse invalid JSON metadata with pytest.raises(json.JSONDecodeError): - await _read_request_body(mock_request) + await read_request_body(mock_request) @pytest.mark.asyncio @@ -320,7 +320,7 @@ async def test_form_data_without_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify all fields are preserved as-is assert result == test_data @@ -351,7 +351,7 @@ async def test_form_data_with_empty_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the metadata was parsed to an empty dict assert "metadata" in result @@ -386,7 +386,7 @@ async def test_form_data_with_dict_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the metadata remains as a dict and is not parsed assert "metadata" in result @@ -417,7 +417,7 @@ async def test_form_data_with_none_metadata(): mock_request.state._cached_headers = None # Parse the form data - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify the metadata remains None (not parsed) assert "metadata" in result @@ -437,7 +437,7 @@ async def test_empty_request_body(): mock_request.scope = {} # Parse the empty body - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify an empty dict is returned assert result == {} @@ -466,7 +466,7 @@ async def test_circular_reference_handling(): mock_request.scope = {} # First parse - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) # Verify initial parse assert result["model"] == "gpt-4" @@ -481,7 +481,7 @@ async def test_circular_reference_handling(): } # Second parse using the same request - will use the modified cached value - result2 = await _read_request_body(mock_request) + result2 = await read_request_body(mock_request) assert "proxy_server_request" not in result2 # This will pass, showing the cache pollution @@ -513,7 +513,7 @@ async def test_json_parsing_error_handling(): # Should raise ProxyException for trailing comma with pytest.raises(ProxyException) as exc_info: - await _read_request_body(mock_request) + await read_request_body(mock_request) assert exc_info.value.code == "400" assert "Invalid JSON payload" in exc_info.value.message @@ -538,7 +538,7 @@ async def test_json_parsing_error_handling(): # Should raise ProxyException for unquoted property with pytest.raises(ProxyException) as exc_info2: - await _read_request_body(mock_request2) + await read_request_body(mock_request2) assert exc_info2.value.code == "400" assert "Invalid JSON payload" in exc_info2.value.message @@ -564,7 +564,7 @@ async def test_json_parsing_error_handling(): mock_request3.scope = {} # Should parse successfully - result = await _read_request_body(mock_request3) + result = await read_request_body(mock_request3) assert result["model"] == "gpt-4o" assert result["input"] == "Run available tools" assert len(result["tools"]) == 1 @@ -597,21 +597,21 @@ async def test_surrogate_repair_skipped_above_size_limit(monkeypatch): small_body = b'{"model":"gpt-4o","x":NaN}' assert len(small_body) <= 100 - repaired = await _read_request_body(_make_json_request(small_body)) + repaired = await read_request_body(_make_json_request(small_body)) assert repaired["model"] == "gpt-4o" padding = "a" * 200 large_body = b'{"model":"gpt-4o","pad":"' + padding.encode() + b'","x":NaN}' assert len(large_body) > 100 with pytest.raises(ProxyException) as exc_info: - await _read_request_body(_make_json_request(large_body)) + await read_request_body(_make_json_request(large_body)) assert exc_info.value.code == "400" assert "Invalid JSON payload" in exc_info.value.message # Disabling the cap (0) restores repair for the same large body, proving the cap # — not the malformed content — is what short-circuits the repair. monkeypatch.setattr(http_parsing_utils, "MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB", 0) - repaired_large = await _read_request_body(_make_json_request(large_body)) + repaired_large = await read_request_body(_make_json_request(large_body)) assert repaired_large["model"] == "gpt-4o" @@ -632,13 +632,13 @@ async def test_lone_surrogate_escape_is_rejected_with_400(content: bytes): """ body = b'{"model":"gpt-4o","messages":[{"role":"user","content":"' + content + b'"}]}' with pytest.raises(ProxyException) as exc_info: - await _read_request_body(_make_json_request(body)) + await read_request_body(_make_json_request(body)) assert exc_info.value.code == "400" assert exc_info.value.type == "invalid_request_error" assert "Invalid JSON payload" in exc_info.value.message paired = body.replace(content, b"say ok \\ud83d\\ude00") - parsed = await _read_request_body(_make_json_request(paired)) + parsed = await read_request_body(_make_json_request(paired)) assert parsed["messages"][0]["content"] == "say ok \U0001f600" @@ -646,7 +646,7 @@ async def test_lone_surrogate_escape_is_rejected_with_400(content: bytes): @pytest.mark.parametrize("media_type", ["application/x-protobuf", "application/protobuf", "application/octet-stream"]) async def test_json_body_under_a_binary_content_type_is_still_parsed(media_type: str): request = _starlette_request(b'{"model": "claude-sonnet-5"}', media_type) - assert await _read_request_body(request) == {"model": "claude-sonnet-5"} + assert await read_request_body(request) == {"model": "claude-sonnet-5"} @pytest.mark.asyncio @@ -920,7 +920,7 @@ async def test_request_body_with_html_script_tags(): mock_request.headers = {"content-type": "application/json"} mock_request.scope = {} - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) assert result["model"] == "gpt-4o" assert len(result["messages"]) == 3 @@ -943,7 +943,7 @@ def test_safe_get_request_headers_caches_on_request_state(): mock_request.state = MagicMock(spec=[]) # empty spec so getattr returns default # First call — should create and cache - result1 = _safe_get_request_headers(mock_request) + result1 = safe_get_request_headers(mock_request) assert result1 == { "content-type": "application/json", "authorization": "Bearer sk-123", @@ -951,7 +951,7 @@ def test_safe_get_request_headers_caches_on_request_state(): assert mock_request.state._cached_headers is result1 # Second call — should return the cached object (same identity) - result2 = _safe_get_request_headers(mock_request) + result2 = safe_get_request_headers(mock_request) assert result2 is result1 @@ -959,7 +959,7 @@ def test_safe_get_request_headers_none_request(): """ Test that _safe_get_request_headers returns empty dict for None request. """ - result = _safe_get_request_headers(None) + result = safe_get_request_headers(None) assert result == {} @@ -971,15 +971,15 @@ def test_safe_get_request_headers_copy_protects_cache(): mock_request.headers = {"authorization": "Bearer sk-123", "host": "localhost"} mock_request.state = MagicMock(spec=[]) - original = _safe_get_request_headers(mock_request) + original = safe_get_request_headers(mock_request) # Simulate what mutation call sites do: copy then pop - mutable = _safe_get_request_headers(mock_request).copy() + mutable = safe_get_request_headers(mock_request).copy() mutable.pop("authorization", None) # Cache must be unaffected - assert "authorization" in _safe_get_request_headers(mock_request) - assert _safe_get_request_headers(mock_request) is original + assert "authorization" in safe_get_request_headers(mock_request) + assert safe_get_request_headers(mock_request) is original def test_safe_get_request_headers_state_unavailable(): @@ -1001,7 +1001,7 @@ def test_safe_get_request_headers_state_unavailable(): mock_request.headers = {"content-type": "application/json"} mock_request.state = ReadOnlyState() - result = _safe_get_request_headers(mock_request) + result = safe_get_request_headers(mock_request) assert result == {"content-type": "application/json"} @@ -1095,7 +1095,7 @@ class TestReadRequestBodyNonCanonicalContentType: mock_request.headers = {"content-type": content_type} mock_request.scope = {} - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) assert result == payload mock_request.form.assert_not_called() @@ -1107,7 +1107,7 @@ class TestReadRequestBodyNonCanonicalContentType: mock_request.headers = {"content-type": "application/x-www-form-urlencoded"} mock_request.scope = {} - result = await _read_request_body(mock_request) + result = await read_request_body(mock_request) assert result == {"k": "v"} mock_request.form.assert_awaited_once() @@ -1136,7 +1136,7 @@ class TestReadRequestBodyFormParseFailure: mock_request.scope = {} with pytest.raises(ProxyException) as exc_info: - await _read_request_body(mock_request) + await read_request_body(mock_request) assert str(exc_info.value.code) == "400" @@ -1341,7 +1341,7 @@ async def test_only_trace_ingest_skips_json_body(method: str, path: str, skip_pa receive, ) - parsed: Final = await _read_request_body(request) + parsed: Final = await read_request_body(request) if skip_parse: assert parsed == {} receive.assert_not_awaited() @@ -1383,7 +1383,7 @@ async def test_otlp_auth_does_not_consume_chunked_bodies_before_the_receiver_lim }, receive, ) - assert await _read_request_body(request) == {} + assert await read_request_body(request) == {} assert received == [] storage = MagicMock() storage.ingest = AsyncMock() @@ -1535,27 +1535,27 @@ def _request_with_body(body: bytes) -> Request_http_parsing: @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_read_request_body_valid_json(): - result = await _read_request_body(_request_with_body(b'{"key": "value"}')) + result = await read_request_body(_request_with_body(b'{"key": "value"}')) assert result == {"key": "value"} @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_read_request_body_empty_body(): - result = await _read_request_body(_request_with_body(b"")) + result = await read_request_body(_request_with_body(b"")) assert result == {} @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_read_request_body_invalid_json(): with pytest.raises(ProxyException): - await _read_request_body(_request_with_body(b'{"key": value}')) + await read_request_body(_request_with_body(b'{"key": value}')) @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_read_request_body_large_payload(): large_payload = '{"key":' + '"a"' * 10**6 + "}" with pytest.raises(ProxyException): - await _read_request_body(_request_with_body(large_payload.encode())) + await read_request_body(_request_with_body(large_payload.encode())) @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio @@ -1563,5 +1563,5 @@ async def test_read_request_body_unexpected_error(): async def receive() -> Message: raise ValueError("Unexpected error") - result = await _read_request_body(_request(receive)) + result = await read_request_body(_request(receive)) assert result == {} diff --git a/tests/unit/proxy/common_utils/test_rbac_utils.py b/tests/unit/proxy/common_utils/test_rbac_utils.py index 54b465a94ba..9c2af565efb 100644 --- a/tests/unit/proxy/common_utils/test_rbac_utils.py +++ b/tests/unit/proxy/common_utils/test_rbac_utils.py @@ -60,9 +60,7 @@ async def test_feature_not_disabled_allows_internal_user(): @pytest.mark.asyncio async def test_feature_not_disabled_allows_vector_stores(): user = _make_user(LitellmUserRoles.INTERNAL_USER.value) - with patch.dict( - _GS_PATH, {"disable_vector_stores_for_internal_users": False}, clear=True - ): + with patch.dict(_GS_PATH, {"disable_vector_stores_for_internal_users": False}, clear=True): await check_feature_access_for_user(user, "vector_stores") @@ -120,7 +118,7 @@ async def test_agents_disabled_team_admin_allowed(): clear=True, ): with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_privileges", + "litellm.proxy.management_endpoints.common_utils.user_has_admin_privileges", new=AsyncMock(return_value=True), ): await check_feature_access_for_user(user, "agents") @@ -138,7 +136,7 @@ async def test_agents_disabled_non_team_admin_blocked(): clear=True, ): with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_privileges", + "litellm.proxy.management_endpoints.common_utils.user_has_admin_privileges", new=AsyncMock(return_value=False), ): with pytest.raises(HTTPException) as exc_info: @@ -158,7 +156,7 @@ async def test_vector_stores_disabled_team_admin_allowed(): clear=True, ): with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_privileges", + "litellm.proxy.management_endpoints.common_utils.user_has_admin_privileges", new=AsyncMock(return_value=True), ): await check_feature_access_for_user(user, "vector_stores") @@ -176,7 +174,7 @@ async def test_vector_stores_disabled_non_team_admin_blocked(): clear=True, ): with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_privileges", + "litellm.proxy.management_endpoints.common_utils.user_has_admin_privileges", new=AsyncMock(return_value=False), ): with pytest.raises(HTTPException) as exc_info: @@ -238,8 +236,6 @@ async def test_org_admin_role_enum_and_string_both_blocked(): with pytest.raises(HTTPException): await check_org_admin_can_generate_keys(user_str) - user_enum = UserAPIKeyAuth( - user_role=LitellmUserRoles.ORG_ADMIN, user_id="user-1" - ) + user_enum = UserAPIKeyAuth(user_role=LitellmUserRoles.ORG_ADMIN, user_id="user-1") with pytest.raises(HTTPException): await check_org_admin_can_generate_keys(user_enum) diff --git a/tests/unit/proxy/common_utils/test_realtime_cache.py b/tests/unit/proxy/common_utils/test_realtime_cache.py index 8316ed1d29a..6f5751fa923 100644 --- a/tests/unit/proxy/common_utils/test_realtime_cache.py +++ b/tests/unit/proxy/common_utils/test_realtime_cache.py @@ -2,21 +2,21 @@ from typing import Any, cast import pytest -from litellm.proxy.common_utils.realtime_utils import _realtime_request_body +from litellm.proxy.common_utils.realtime_utils import realtime_request_body from litellm.proxy.proxy_server import _realtime_query_params_template @pytest.fixture(autouse=True) def clear_realtime_caches(): - _realtime_request_body.cache_clear() + realtime_request_body.cache_clear() _realtime_query_params_template.cache_clear() yield - _realtime_request_body.cache_clear() + realtime_request_body.cache_clear() _realtime_query_params_template.cache_clear() def test_realtime_request_body_returns_immutable_bytes(): - cached_body = _realtime_request_body("gpt-4o") + cached_body = realtime_request_body("gpt-4o") with pytest.raises(TypeError): cast(Any, cached_body)[0] = ord("x") @@ -30,9 +30,9 @@ def test_realtime_query_params_template_returns_immutable_tuples(): def test_realtime_request_body_caches_each_model_separately(): - gpt4o_body_first = _realtime_request_body("gpt-4o") - gpt4o_body_second = _realtime_request_body("gpt-4o") - gpt4o_mini_body = _realtime_request_body("gpt-4o-mini") + gpt4o_body_first = realtime_request_body("gpt-4o") + gpt4o_body_second = realtime_request_body("gpt-4o") + gpt4o_mini_body = realtime_request_body("gpt-4o-mini") assert gpt4o_body_first is gpt4o_body_second assert gpt4o_body_first == b'{"model": "gpt-4o"}' @@ -44,9 +44,7 @@ def test_realtime_query_params_template_caches_each_pair_separately(): params_with_intent_first = _realtime_query_params_template("gpt-4o", "intent-a") params_with_intent_second = _realtime_query_params_template("gpt-4o", "intent-a") params_without_intent = _realtime_query_params_template("gpt-4o", None) - params_transcription_without_model = _realtime_query_params_template( - None, "transcription" - ) + params_transcription_without_model = _realtime_query_params_template(None, "transcription") assert params_with_intent_first is params_with_intent_second assert params_with_intent_first == (("model", "gpt-4o"), ("intent", "intent-a")) diff --git a/tests/unit/proxy/common_utils/test_upsert_budget_membership.py b/tests/unit/proxy/common_utils/test_upsert_budget_membership.py index 7a75ec395f1..74d1748d4de 100644 --- a/tests/unit/proxy/common_utils/test_upsert_budget_membership.py +++ b/tests/unit/proxy/common_utils/test_upsert_budget_membership.py @@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from litellm.proxy.management_endpoints.common_utils import ( - _upsert_budget_and_membership, + upsert_budget_and_membership, ) # --------------------------------------------------------------------------- @@ -68,7 +68,7 @@ def stored_budget_row(mock_tx): # role must not silently wipe their budget. @pytest.mark.asyncio async def test_empty_patch_is_noop(mock_tx, fake_user): - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-1", user_id="user-1", @@ -89,7 +89,7 @@ async def test_empty_patch_is_noop(mock_tx, fake_user): async def test_clearing_all_limits_disconnects(mock_tx, fake_user): mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=100.0)) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-1", user_id="user-1", @@ -114,7 +114,7 @@ async def test_clear_one_field_keeps_others(mock_tx, fake_user): return_value=budget_row(max_budget=100.0, budget_duration="24h") ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-1", user_id="user-1", @@ -140,7 +140,7 @@ async def test_clear_one_field_keeps_others(mock_tx, fake_user): async def test_update_in_place_seeds_reset_at(mock_tx, fake_user): mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=20.0)) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-dur", user_id="user-dur", @@ -165,7 +165,7 @@ async def test_update_in_place_seeds_reset_at(mock_tx, fake_user): async def test_update_in_place_single_field_leaves_reset_at_alone(mock_tx, fake_user): mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=50.0)) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-rpm", user_id="user-rpm", @@ -185,7 +185,7 @@ async def test_update_in_place_single_field_leaves_reset_at_alone(mock_tx, fake_ # the duration and a future reset time, then links the membership. @pytest.mark.asyncio async def test_create_seeds_reset_at_and_links(mock_tx, fake_user): - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-new", user_id="user-new", @@ -220,7 +220,7 @@ async def test_create_seeds_reset_at_and_links(mock_tx, fake_user): @pytest.mark.asyncio async def test_create_from_temp_budget_pair_only(mock_tx, fake_user): expiry = datetime(2100, 1, 1, tzinfo=timezone.utc) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-new", user_id="user-new", @@ -241,7 +241,7 @@ async def test_create_from_temp_pair_never_snapshots_team_default(mock_tx, fake_ mock_tx.litellm_budgettable.find_unique = AsyncMock( return_value=budget_row(budget_id="team-default-budget-1", max_budget=0.4, rpm_limit=10) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-unlinked", @@ -262,7 +262,7 @@ async def test_temp_pair_on_shared_default_member_creates_bare_row(mock_tx, fake mock_tx.litellm_budgettable.find_unique = AsyncMock( return_value=budget_row(budget_id="team-default-budget-1", max_budget=0.4, rpm_limit=10) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-on-default", @@ -280,7 +280,7 @@ async def test_temp_pair_on_shared_default_member_creates_bare_row(mock_tx, fake @pytest.mark.asyncio async def test_clearing_temp_pair_on_shared_default_member_is_noop(mock_tx, fake_user): - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-on-default", @@ -302,7 +302,7 @@ async def test_temp_pair_with_permanent_field_still_clones_shared_default(mock_t mock_tx.litellm_budgettable.find_unique = AsyncMock( return_value=budget_row(budget_id="team-default-budget-1", max_budget=0.4, rpm_limit=10) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-on-default", @@ -326,7 +326,7 @@ async def test_create_from_plain_patch_does_not_snapshot_team_default(mock_tx, f mock_tx.litellm_budgettable.find_unique = AsyncMock( return_value=budget_row(budget_id="team-default-budget-1", max_budget=0.4, rpm_limit=10) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-default", user_id="user-unlinked", @@ -363,7 +363,7 @@ async def test_clone_on_write_from_shared_default(mock_tx, fake_user): ) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-shared", user_id="user-shared", @@ -416,7 +416,7 @@ async def test_clone_on_write_clears_duration(mock_tx, fake_user): ) ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-shared", user_id="user-shared", @@ -444,7 +444,7 @@ async def test_clone_on_write_clears_duration(mock_tx, fake_user): async def test_private_budget_updates_in_place(mock_tx, fake_user): mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row(max_budget=10.0)) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( mock_tx, team_id="team-mixed", user_id="user-private", diff --git a/tests/unit/proxy/db/db_transaction_queue/test_redis_update_buffer.py b/tests/unit/proxy/db/db_transaction_queue/test_redis_update_buffer.py index e04e2402e1b..5ecc8784384 100644 --- a/tests/unit/proxy/db/db_transaction_queue/test_redis_update_buffer.py +++ b/tests/unit/proxy/db/db_transaction_queue/test_redis_update_buffer.py @@ -675,7 +675,7 @@ class _ListRedis: async def test_store_spend_logs_in_redis_drops_oldest_rows_past_the_cap(): redis = _ListRedis() buffer = RedisUpdateBuffer(redis_cache=redis) - buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=True) + buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=True) assert await buffer.store_spend_logs_in_redis([{"request_id": "old"}, {"request_id": "mid"}], max_rows=2) is True assert await buffer.store_spend_logs_in_redis([{"request_id": "new"}], max_rows=2) is True @@ -697,7 +697,7 @@ async def test_store_spend_logs_in_redis_reports_failure_without_redis(): async def test_store_spend_logs_in_redis_is_off_unless_transaction_buffering_is_enabled(): redis = _ListRedis() buffer = RedisUpdateBuffer(redis_cache=redis) - buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=False) + buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=False) assert await buffer.store_spend_logs_in_redis([{"request_id": "a"}]) is False assert redis.rows == [] diff --git a/tests/unit/proxy/db/test_db_spend_update_writer.py b/tests/unit/proxy/db/test_db_spend_update_writer.py index 8822ecd5bd0..dc9b8fd395b 100644 --- a/tests/unit/proxy/db/test_db_spend_update_writer.py +++ b/tests/unit/proxy/db/test_db_spend_update_writer.py @@ -997,7 +997,7 @@ async def test_commit_spend_updates_to_db_increments_agent_spend(): "agent_list_transactions": {agent_id: response_cost}, } - with patch("litellm.proxy.utils._raise_failed_update_spend_exception"): + with patch("litellm.proxy.utils.raise_failed_update_spend_exception"): await db_writer._commit_spend_updates_to_db( prisma_client=mock_prisma_client, n_retry_times=0, @@ -2029,7 +2029,7 @@ async def test_commit_key_spend_updates_includes_last_active(): before_call = datetime.now(timezone.utc) - with patch("litellm.proxy.utils._raise_failed_update_spend_exception"): + with patch("litellm.proxy.utils.raise_failed_update_spend_exception"): await db_writer._commit_spend_updates_to_db( prisma_client=mock_prisma_client, n_retry_times=0, @@ -3608,7 +3608,7 @@ async def test_commit_spend_updates_to_db_does_not_stamp_key_settings_updated_at "agent_list_transactions": {}, } - with patch("litellm.proxy.utils._raise_failed_update_spend_exception"): + with patch("litellm.proxy.utils.raise_failed_update_spend_exception"): await db_writer._commit_spend_updates_to_db( prisma_client=mock_prisma_client, n_retry_times=0, diff --git a/tests/unit/proxy/db/test_update_daily_tag_spend.py b/tests/unit/proxy/db/test_update_daily_tag_spend.py index dda41ab4543..d720f607942 100644 --- a/tests/unit/proxy/db/test_update_daily_tag_spend.py +++ b/tests/unit/proxy/db/test_update_daily_tag_spend.py @@ -14,23 +14,23 @@ async def test_update_daily_tag_spend_delegates_to_tag_commit_writer(): prisma_client = MagicMock() proxy_logging_obj = MagicMock() redis_update_buffer = MagicMock() - redis_update_buffer._should_commit_spend_updates_to_redis.return_value = False + redis_update_buffer.should_commit_spend_updates_to_redis.return_value = False proxy_logging_obj.db_spend_update_writer = MagicMock() proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock() - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() await update_daily_tag_spend( prisma_client, proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db.assert_awaited_once_with( + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db.assert_awaited_once_with( prisma_client=prisma_client, n_retry_times=3, proxy_logging_obj=proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis.assert_not_awaited() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis.assert_not_awaited() @pytest.mark.asyncio @@ -38,11 +38,11 @@ async def test_update_daily_tag_spend_logs_error_and_does_not_raise(): prisma_client = MagicMock() proxy_logging_obj = MagicMock() redis_update_buffer = MagicMock() - redis_update_buffer._should_commit_spend_updates_to_redis.return_value = False + redis_update_buffer.should_commit_spend_updates_to_redis.return_value = False proxy_logging_obj.db_spend_update_writer = MagicMock() proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock(side_effect=ValueError("boom")) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db = AsyncMock(side_effect=ValueError("boom")) + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() with patch("litellm.proxy.utils.verbose_proxy_logger.error") as error_logger: await update_daily_tag_spend( @@ -50,7 +50,7 @@ async def test_update_daily_tag_spend_logs_error_and_does_not_raise(): proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db.assert_awaited_once() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db.assert_awaited_once() error_logger.assert_called_once() @@ -59,23 +59,23 @@ async def test_update_daily_tag_spend_uses_redis_writer_when_enabled(): prisma_client = MagicMock() proxy_logging_obj = MagicMock() redis_update_buffer = MagicMock() - redis_update_buffer._should_commit_spend_updates_to_redis.return_value = True + redis_update_buffer.should_commit_spend_updates_to_redis.return_value = True proxy_logging_obj.db_spend_update_writer = MagicMock() - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db = AsyncMock() proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() await update_daily_tag_spend( prisma_client, proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis.assert_awaited_once_with( + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db_with_redis.assert_awaited_once_with( prisma_client=prisma_client, n_retry_times=3, proxy_logging_obj=proxy_logging_obj, ) - proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db.assert_not_awaited() + proxy_logging_obj.db_spend_update_writer.commit_daily_tag_spend_to_db.assert_not_awaited() @pytest.mark.asyncio diff --git a/tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py b/tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py index cddb0e526b4..79062f15eca 100644 --- a/tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py +++ b/tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py @@ -63,9 +63,7 @@ class TestMergeQueryParamsIntoData: assert "api_key" not in data def test_litellm_params_template_json_is_expanded(self): - template = json.dumps( - {"api_key": "AIzaFromTemplate", "api_base": "https://example.com"} - ) + template = json.dumps({"api_key": "AIzaFromTemplate", "api_base": "https://example.com"}) from urllib.parse import quote request = _make_request(f"litellm_params_template={quote(template)}") @@ -77,9 +75,7 @@ class TestMergeQueryParamsIntoData: assert "litellm_params_template" not in data def test_litellm_params_template_does_not_overwrite_existing(self): - template = json.dumps( - {"api_key": "FromTemplate", "custom_llm_provider": "openai"} - ) + template = json.dumps({"api_key": "FromTemplate", "custom_llm_provider": "openai"}) from urllib.parse import quote request = _make_request(f"litellm_params_template={quote(template)}") @@ -160,18 +156,14 @@ def _make_endpoint_request(query_string: str = "") -> MagicMock: @pytest.mark.asyncio -async def test_list_gemini_agents_passes_api_key_to_processor( - mock_srv, user_api_key_dict -): +async def test_list_gemini_agents_passes_api_key_to_processor(mock_srv, user_api_key_dict): from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents template = json.dumps({"api_key": "AIzaListTest"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -188,18 +180,14 @@ async def test_list_gemini_agents_passes_api_key_to_processor( @pytest.mark.asyncio -async def test_get_gemini_agent_passes_api_key_to_processor( - mock_srv, user_api_key_dict -): +async def test_get_gemini_agent_passes_api_key_to_processor(mock_srv, user_api_key_dict): from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import get_gemini_agent template = json.dumps({"api_key": "AIzaGetTest"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -218,18 +206,14 @@ async def test_get_gemini_agent_passes_api_key_to_processor( @pytest.mark.asyncio -async def test_delete_gemini_agent_passes_api_key_to_processor( - mock_srv, user_api_key_dict -): +async def test_delete_gemini_agent_passes_api_key_to_processor(mock_srv, user_api_key_dict): from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import delete_gemini_agent template = json.dumps({"api_key": "AIzaDeleteTest"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -248,9 +232,7 @@ async def test_delete_gemini_agent_passes_api_key_to_processor( @pytest.mark.asyncio -async def test_list_gemini_agent_versions_passes_api_key_to_processor( - mock_srv, user_api_key_dict -): +async def test_list_gemini_agent_versions_passes_api_key_to_processor(mock_srv, user_api_key_dict): from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import ( @@ -259,9 +241,7 @@ async def test_list_gemini_agent_versions_passes_api_key_to_processor( template = json.dumps({"api_key": "AIzaVersionsTest"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -280,17 +260,13 @@ async def test_list_gemini_agent_versions_passes_api_key_to_processor( @pytest.mark.asyncio -async def test_get_gemini_agent_name_not_overwritten_by_query_param( - mock_srv, user_api_key_dict -): +async def test_get_gemini_agent_name_not_overwritten_by_query_param(mock_srv, user_api_key_dict): """Path-param ``name`` must not be replaced by an attacker-controlled query param.""" from urllib.parse import quote from litellm.proxy.google_endpoints.agents_endpoints import get_gemini_agent - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -299,9 +275,7 @@ async def test_get_gemini_agent_name_not_overwritten_by_query_param( # ``api_key`` is supplied via the JSON template (required for non-admin # callers — see test_*_non_admin_without_api_key_is_rejected below). template = json.dumps({"api_key": "AIzaTest"}) - request = _make_endpoint_request( - f"name=INJECTED&litellm_params_template={quote(template)}" - ) + request = _make_endpoint_request(f"name=INJECTED&litellm_params_template={quote(template)}") await get_gemini_agent( request=request, name="real-agent", @@ -321,9 +295,7 @@ async def test_list_agents_template_via_query_param(mock_srv, user_api_key_dict) template = json.dumps({"api_key": "TemplateKey", "vertex_project": "proj-x"}) - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -356,9 +328,7 @@ def proxy_admin_user_api_key_dict(): @pytest.mark.asyncio -async def test_list_agents_non_admin_without_api_key_is_rejected( - mock_srv, user_api_key_dict -): +async def test_list_agents_non_admin_without_api_key_is_rejected(mock_srv, user_api_key_dict): """Non-admin callers must supply an explicit api_key — the proxy must not silently fall back to the operator's shared GOOGLE_API_KEY/GEMINI_API_KEY. """ @@ -366,9 +336,7 @@ async def test_list_agents_non_admin_without_api_key_is_rejected( from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -385,16 +353,12 @@ async def test_list_agents_non_admin_without_api_key_is_rejected( @pytest.mark.asyncio -async def test_delete_agent_non_admin_without_api_key_is_rejected( - mock_srv, user_api_key_dict -): +async def test_delete_agent_non_admin_without_api_key_is_rejected(mock_srv, user_api_key_dict): from fastapi import HTTPException from litellm.proxy.google_endpoints.agents_endpoints import delete_gemini_agent - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) @@ -411,19 +375,15 @@ async def test_delete_agent_non_admin_without_api_key_is_rejected( @pytest.mark.asyncio -async def test_create_agent_non_admin_without_api_key_is_rejected( - mock_srv, user_api_key_dict -): +async def test_create_agent_non_admin_without_api_key_is_rejected(mock_srv, user_api_key_dict): from fastapi import HTTPException from litellm.proxy.google_endpoints.agents_endpoints import create_gemini_agent with ( + patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor, patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor, - patch( - "litellm.proxy.google_endpoints.agents_endpoints._read_request_body", + "litellm.proxy.google_endpoints.agents_endpoints.read_request_body", new=AsyncMock(return_value={"name": "agent-1", "base_agent": "waverunner"}), ), ): @@ -442,15 +402,11 @@ async def test_create_agent_non_admin_without_api_key_is_rejected( @pytest.mark.asyncio -async def test_list_agents_proxy_admin_may_use_env_fallback( - mock_srv, proxy_admin_user_api_key_dict -): +async def test_list_agents_proxy_admin_may_use_env_fallback(mock_srv, proxy_admin_user_api_key_dict): """Proxy admins (master key) keep the env-fallback convenience.""" from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents - with patch( - "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing") as MockProcessor: instance = MockProcessor.return_value instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma.py b/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma.py index e23796705b3..f06ba3ee68e 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma.py @@ -11,7 +11,7 @@ import litellm from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.guardrails.guardrail_hooks.presidio import _OPTIONAL_PresidioPIIMasking +from litellm.proxy.guardrails.guardrail_hooks.presidio import OPTIONAL_PresidioPIIMasking from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import StandardLoggingPayload from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome @@ -34,7 +34,7 @@ async def test_standard_logging_payload_includes_guardrail_information(): """ test_custom_logger = CustomLoggerForTesting() litellm.callbacks = [test_custom_logger] - presidio_guard = _OPTIONAL_PresidioPIIMasking( + presidio_guard = OPTIONAL_PresidioPIIMasking( guardrail_name="presidio_guard", event_hook=GuardrailEventHooks.pre_call, presidio_analyzer_api_base="https://mock-presidio-analyzer.com/", diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index d0b0799817b..01be28470d1 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -3520,7 +3520,7 @@ async def test_chat_completion_modify_response_exception_streaming_logging_obj_n raise exc with ( - patch("litellm.proxy.proxy_server._read_request_body", AsyncMock(return_value=request_data)), + patch("litellm.proxy.proxy_server.read_request_body", AsyncMock(return_value=request_data)), patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), patch( "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request", diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py index ff035a3741e..0cac6228085 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -23,7 +23,7 @@ from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.presidio import ( PresidioPerRequestConfig, - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) from litellm.exceptions import GuardrailRaisedException from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType @@ -80,7 +80,7 @@ def _make_mock_session_iterator(json_response, status=200, content_type="applica @pytest.fixture def presidio_guardrail(): """Create a Presidio guardrail instance for testing""" - return _OPTIONAL_PresidioPIIMasking( + return OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=False, pii_entities_config={ @@ -638,7 +638,7 @@ async def test_logging_only_does_not_mask_pre_call_request(mock_user_api_key, mo causing the model's response to contain anonymization tokens (e.g. ) instead of the real output. """ - presidio_guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, logging_only=True, pii_entities_config={PiiEntityType.PHONE_NUMBER: PiiAction.MASK}, @@ -677,7 +677,7 @@ async def test_presidio_sets_guardrail_information_in_request_data(): This validates that add_standard_logging_guardrail_information_to_request_data correctly sets the guardrail information that will be used for logging. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", output_parse_pii=True, mock_testing=True, @@ -735,7 +735,7 @@ async def test_request_data_flows_to_apply_guardrail(): This validates the fix where guardrail translation handler passes data as request_data to apply_guardrail so guardrails can store metadata for logging. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", output_parse_pii=True, mock_testing=True, @@ -775,7 +775,7 @@ async def test_output_masking_apply_to_output_only(mock_user_api_key): Ensure output masking runs when apply_to_output is enabled. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, pii_entities_config={PiiEntityType.CREDIT_CARD: PiiAction.MASK}, @@ -842,8 +842,8 @@ async def test_presidio_filter_scope_initializer(monkeypatch): import litellm.proxy.guardrails.guardrail_hooks.presidio as presidio_mod import litellm.proxy.guardrails.guardrail_initializers as gi - monkeypatch.setattr(presidio_mod, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False) - monkeypatch.setattr(gi, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False) + monkeypatch.setattr(presidio_mod, "OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False) + monkeypatch.setattr(gi, "OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False) # input-only created.clear() @@ -982,7 +982,7 @@ async def test_analyze_text_with_empty_string(): Should return empty list without making API call to Presidio. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test:5002/", presidio_anonymizer_api_base="http://test:5001/", output_parse_pii=False, @@ -1015,7 +1015,7 @@ async def test_analyze_text_error_dict_handling(): When Presidio returns {'error': 'No text provided'}, should handle gracefully instead of crashing with TypeError. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1044,7 +1044,7 @@ async def test_analyze_text_string_response_handling(): When Presidio returns a string (e.g. error message from websearch/hosted models), should handle gracefully instead of crashing with TypeError about mapping vs str. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1069,7 +1069,7 @@ async def test_analyze_text_invalid_response_raises_when_block_configured(): When pii_entities_config has BLOCK and Presidio returns invalid response, should raise GuardrailRaisedException (fail-closed) rather than silently allowing content. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1096,7 +1096,7 @@ async def test_analyze_text_invalid_response_raises_when_mask_configured(): When pii_entities_config has MASK and Presidio returns invalid response, should raise GuardrailRaisedException (fail-closed) because PII masking is expected. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1125,7 +1125,7 @@ async def test_analyze_text_list_with_non_dict_items(): When Presidio returns a list containing strings (malformed response), should skip invalid items and return parsed valid ones. """ - presidio = _OPTIONAL_PresidioPIIMasking( + presidio = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://mock-presidio:5002/", presidio_anonymizer_api_base="http://mock-presidio:5001/", output_parse_pii=False, @@ -1210,7 +1210,7 @@ def test_filter_drops_low_score_detection(): """ Detections below the configured score threshold should be removed. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, ) @@ -1224,7 +1224,7 @@ def test_filter_preserves_high_score_detection(): """ Detections meeting the score threshold should be preserved. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, ) @@ -1239,7 +1239,7 @@ def test_no_thresholds_returns_all(): """ With no thresholds configured, all detections are kept. """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True) analyze_results = [ {"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.1, "start": 0, "end": 4}, { @@ -1258,7 +1258,7 @@ def test_entity_specific_threshold_only_applies_to_that_entity(): """ Entity-specific thresholds do not affect other entity types. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, ) @@ -1282,7 +1282,7 @@ def test_filter_uses_default_all_threshold(): """ Default ALL threshold applies to any entity without a specific override. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={"ALL": 0.75}, ) @@ -1305,7 +1305,7 @@ def test_entity_specific_overrides_default_threshold(): """ Entity-specific threshold should override the ALL default. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={ "ALL": 0.8, @@ -1333,7 +1333,7 @@ async def test_anonymize_skips_when_no_detections_after_filter(): """ When all detections are filtered out, anonymize_text should return the original text. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, ) @@ -1359,7 +1359,7 @@ def test_blocking_respects_threshold_filter(): """ Entities filtered out by score should not trigger blocking, but high-score detections should. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, pii_entities_config={PiiEntityType.CREDIT_CARD: PiiAction.BLOCK}, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.9}, @@ -1379,7 +1379,7 @@ def test_update_in_memory_applies_score_thresholds(): """ update_in_memory_litellm_params should refresh score thresholds. """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True) assert guardrail.presidio_score_thresholds == {} params = LitellmParams( @@ -1458,7 +1458,7 @@ async def test_streaming_with_bytes_chunks_does_not_crash(mock_user_api_key): gracefully handle raw bytes in the stream instead of crashing with 'bytes' object has no attribute 'id'. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": "redacted"}, @@ -1491,7 +1491,7 @@ def test_entity_deny_list_filters_detections(): """ Verify presidio_entities_deny_list removes matching entity types. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_entities_deny_list=["US_DRIVER_LICENSE"], ) @@ -1511,7 +1511,7 @@ def test_deny_list_and_score_threshold_combined(): """ Verify deny list + score threshold work together correctly. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, presidio_entities_deny_list=["US_DRIVER_LICENSE"], presidio_score_thresholds={"ALL": 0.8}, @@ -1538,7 +1538,7 @@ async def test_analyze_text_non_json_content_type_fail_closed(): Test that analyze_text raises GuardrailRaisedException when Presidio health endpoint returns text/html and fail-closed is enabled. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", pii_entities_config={"PERSON": PiiAction.BLOCK}, @@ -1569,7 +1569,7 @@ async def test_analyze_text_non_json_content_type_fail_open(): Test that analyze_text returns empty list when Presidio returns text/html and fail-closed is NOT enabled. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -1596,7 +1596,7 @@ async def test_analyze_text_http_error_status(): """ Test that analyze_text handles 5xx HTTP errors properly. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", pii_entities_config={"PERSON": PiiAction.BLOCK}, @@ -1625,7 +1625,7 @@ async def test_anonymize_text_non_json_content_type(): """ Test that anonymize_text raises Exception for non-JSON responses. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -1653,7 +1653,7 @@ async def test_anonymize_text_http_error_status(): """ Test that anonymize_text raises Exception on HTTP error. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -1684,7 +1684,7 @@ async def test_pii_tokens_stored_in_metadata_not_top_level(presidio_guardrail): providers like Anthropic, which reject unknown fields with 'pii_tokens: Extra inputs are not permitted'. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, pii_entities_config={ @@ -1743,7 +1743,7 @@ async def test_pii_tokens_in_metadata_used_for_unmasking(): Regression test: _process_response_for_pii must read pii_tokens from data['metadata']['pii_tokens'] and correctly unmask the response. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -1786,7 +1786,7 @@ def test_event_hook_auto_expansion_for_all_string_hooks(initial_hook): 'post_call' to event_hook regardless of the initial string hook value, not just when it's 'pre_call'. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, event_hook=initial_hook, @@ -1798,7 +1798,7 @@ def test_event_hook_auto_expansion_for_all_string_hooks(initial_hook): def test_event_hook_no_expansion_when_already_post_call(): """post_call alone should stay as-is — no expansion needed.""" - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, event_hook="post_call", @@ -1813,7 +1813,7 @@ async def test_metadata_none_does_not_crash(): Regression test: if metadata is explicitly None in request_data, the guardrail must not crash with TypeError on the write or read path. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -1859,7 +1859,7 @@ def test_unmask_exact_match_with_sequential_tokens(): Normal unmasking: LLM echoes numbered tokens verbatim → original PII restored. """ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) pii_tokens = { @@ -1867,7 +1867,7 @@ def test_unmask_exact_match_with_sequential_tokens(): "": "555-123-4567", } text = "Hello , your number is ." - result = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) assert result == "Hello John Smith, your number is 555-123-4567." @@ -1876,7 +1876,7 @@ def test_unmask_multiple_same_entity_type(): Two phone numbers get distinct numbered tokens and unmask correctly. """ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) pii_tokens = { @@ -1884,7 +1884,7 @@ def test_unmask_multiple_same_entity_type(): "": "555-222-0000", } text = "Call or ." - result = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) assert result == "Call 555-111-0000 or 555-222-0000." @@ -1894,7 +1894,7 @@ def test_unmask_graceful_degradation(): in the output — clean and readable, not garbage hex. """ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) pii_tokens = { @@ -1902,7 +1902,7 @@ def test_unmask_graceful_degradation(): } # LLM paraphrased instead of echoing the token text = "I see you provided a name." - result = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_pii_text(text, pii_tokens) # No change — no garbage, just clean text assert result == text @@ -1918,7 +1918,7 @@ async def test_anonymize_text_multiple_items_position_correctness(): Regression test: when multiple PII items exist, coordinates reference the ORIGINAL text. Processing in reverse order prevents coordinate drift. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -1985,7 +1985,7 @@ async def test_anthropic_native_response_unmasking(): Anthropic native dict responses (type='message') should be unmasked when output_parse_pii is enabled. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -2032,7 +2032,7 @@ async def test_anthropic_native_response_masking(): Anthropic native dict responses should be masked when apply_to_output is enabled. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, ) @@ -2071,7 +2071,7 @@ async def test_anthropic_native_response_non_text_blocks_untouched(): Non-text blocks (tool_use, thinking) in Anthropic responses should be left untouched during unmasking. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -2123,7 +2123,7 @@ async def test_streaming_bytes_chunks_are_yielded_not_discarded(): through the streaming hook, not silently discarded. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, ) @@ -2151,7 +2151,7 @@ async def test_streaming_unmask_path_bytes_passthrough(): """ Bytes chunks in the unmasking path should also pass through. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -2183,7 +2183,7 @@ async def test_apply_to_output_streaming_unknown_events_passthrough(): Regression test: /v1/responses-style event objects (neither bytes nor ModelResponseStream) must be preserved in order and not dropped. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, ) @@ -2227,7 +2227,7 @@ async def test_apply_to_output_streaming_mixed_chunks_flushes_and_warns(): a buffered ModelResponseStream chunk followed by unknown responses-style events should be preserved, and masking skip should be visible via warnings. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, ) @@ -2284,7 +2284,7 @@ async def test_apply_guardrail_unmask_on_response(output_parse_pii: bool) -> Non When input_type is 'response' and pii_tokens exist, apply_guardrail should unmask text instead of masking it. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", output_parse_pii=output_parse_pii, mock_testing=True, @@ -2322,7 +2322,7 @@ async def test_standalone_scans_without_restoration_tokens(input_type: Literal[" """ Standalone callbacks retain scanning without tokens, including MCP results. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", event_hook="post_mcp_call", output_parse_pii=True, @@ -2375,7 +2375,7 @@ async def test_apply_to_output_streaming_chat_chunks_are_masked_as_one_response( Structured chat completion chunks are buffered, assembled and masked as a whole, so a card number split across deltas cannot reach the caller. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": "my card is "}, @@ -2406,7 +2406,7 @@ async def test_apply_to_output_streaming_bytes_after_chat_chunks_are_passed_thro Once structured chunks have been buffered, a trailing bytes frame belongs to the same stream and must be forwarded rather than treated as a new SSE stream. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": "hello"}, @@ -2438,7 +2438,7 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_masks_text_split_ac bytes. Output masking must run over the whole content block so a card number split across text_delta events cannot reach the caller. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2488,7 +2488,7 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_masks_text_split_ac @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_sse_bytes_without_pii_are_forwarded_unchanged(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": "Hello world"}, @@ -2538,7 +2538,7 @@ def _gemini_sse(text: str) -> bytes: @pytest.mark.asyncio async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incrementally_until_upstream_aborts(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2567,7 +2567,7 @@ async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incremen @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_first_frame_split_across_transport_chunks_is_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2634,7 +2634,7 @@ def _anthropic_stream_head() -> list[bytes]: @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_first_frame_split_inside_a_utf8_character_is_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2677,7 +2677,7 @@ async def test_apply_to_output_streaming_anthropic_first_frame_split_inside_a_ut @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_stream_led_by_sse_comment_keepalive_is_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2714,7 +2714,7 @@ async def test_apply_to_output_streaming_anthropic_stream_led_by_sse_comment_kee @pytest.mark.asyncio async def test_apply_to_output_streaming_anthropic_stream_led_by_data_less_ping_event_is_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2751,7 +2751,7 @@ async def test_apply_to_output_streaming_anthropic_stream_led_by_data_less_ping_ @pytest.mark.asyncio async def test_apply_to_output_streaming_leading_keepalive_is_forwarded_before_upstream_data_arrives(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2791,7 +2791,7 @@ async def test_apply_to_output_streaming_leading_keepalive_is_forwarded_before_u @pytest.mark.asyncio async def test_apply_to_output_streaming_leading_comments_over_the_frame_cap_are_still_masked(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2832,7 +2832,7 @@ async def test_apply_to_output_streaming_leading_comments_over_the_frame_cap_are @pytest.mark.asyncio async def test_apply_to_output_streaming_comment_only_stream_is_forwarded_unchanged(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2854,7 +2854,7 @@ async def test_apply_to_output_streaming_comment_only_stream_is_forwarded_unchan @pytest.mark.asyncio async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_streams_incrementally(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2885,7 +2885,7 @@ async def test_apply_to_output_streaming_gemini_first_frame_split_across_transpo @pytest.mark.asyncio async def test_apply_to_output_streaming_unterminated_first_frame_is_released_once_it_exceeds_the_cap(): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -2920,7 +2920,7 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_fail_closed_when_pr must surface as an error to the caller: replaying the unscanned frames would hand over whatever PII the model generated. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, presidio_analyzer_api_base="http://127.0.0.1:9", @@ -2970,7 +2970,7 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_block_action_raises A BLOCK on generated PII must refuse the streaming /v1/messages response the same way it refuses the non streaming one, not replay the raw frames. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( apply_to_output=True, mock_testing=False, presidio_analyzer_api_base="http://test-analyzer/", @@ -3023,7 +3023,7 @@ async def test_apply_to_output_streaming_propagates_upstream_error_when_nothing_ An upstream guardrail that rejects the stream before the first chunk must surface as an error to the caller, not as an empty 200 stream. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, @@ -3050,7 +3050,7 @@ async def test_output_parse_pii_streaming_responses_events_passthrough( Regression test: when output_parse_pii=True and pii_tokens exist, /v1/responses streaming events must pass through instead of being dropped. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3099,7 +3099,7 @@ async def test_output_parse_pii_streaming_responses_completed_event_unmasked( ) from litellm.types.responses.main import GenericResponseOutputItem, OutputText - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3159,7 +3159,7 @@ async def test_output_parse_pii_streaming_mixed_chunks_flushes_buffered( chunks must still be forwarded (in order) instead of being dropped at the saw_non_chat_chunk early return. """ - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3244,7 +3244,7 @@ async def test_anonymize_text_uses_correct_positions_no_parse_pii(): ], } - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -3318,7 +3318,7 @@ async def test_anonymize_text_uses_correct_positions_with_parse_pii(): ], } - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", mock_testing=False, @@ -3370,7 +3370,7 @@ def test_unmask_sse_bytes_chunk_replaces_text_delta(): } chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8") - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) decoded = result.decode("utf-8") parsed = json.loads(decoded.split("data: ", 1)[1].strip()) @@ -3385,7 +3385,7 @@ def test_unmask_sse_bytes_chunk_ignores_non_text_delta(): # message_start event — no delta event = {"type": "message_start", "message": {"id": "msg_01", "role": "assistant"}} chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8") - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) assert result == chunk # input_json_delta — should not be touched @@ -3395,19 +3395,19 @@ def test_unmask_sse_bytes_chunk_ignores_non_text_delta(): "delta": {"type": "input_json_delta", "partial_json": '{"name": ""}'}, } chunk2 = ("data: " + json.dumps(event2) + "\n\n").encode("utf-8") - result2 = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk2, pii_tokens) + result2 = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk2, pii_tokens) assert result2 == chunk2 def test_unmask_sse_bytes_chunk_handles_malformed_json(): chunk = b"data: {not valid json}\n\n" - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"": "Bobby"}) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"": "Bobby"}) assert result == chunk def test_unmask_sse_bytes_chunk_handles_unicode_decode_error(): chunk = b"\xff\xfe invalid utf-8" - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"": "Bobby"}) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"": "Bobby"}) assert result == chunk @@ -3422,7 +3422,7 @@ def test_unmask_sse_bytes_chunk_non_ascii_pii_not_escaped(): } chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8") - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) decoded = result.decode("utf-8") assert "Jos\\u" not in decoded @@ -3441,7 +3441,7 @@ def test_unmask_sse_bytes_chunk_handles_crlf_line_endings(): } crlf_chunk = ("data: " + json.dumps(event) + "\r\ndata: [DONE]\r\n").encode("utf-8") - result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(crlf_chunk, pii_tokens) + result = OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(crlf_chunk, pii_tokens) decoded = result.decode("utf-8") parsed = json.loads(decoded.split("data: ", 1)[1].split("\n")[0].strip()) @@ -3453,7 +3453,7 @@ def test_unmask_sse_bytes_chunk_handles_crlf_line_endings(): async def test_stream_pii_unmasking_unmaskes_bytes_chunks(mock_user_api_key): import json - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3489,7 +3489,7 @@ async def test_stream_pii_unmasking_unmaskes_bytes_chunks(mock_user_api_key): @pytest.mark.asyncio async def test_stream_pii_unmasking_passthrough_when_no_tokens(mock_user_api_key): - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, output_parse_pii=True, ) @@ -3514,7 +3514,7 @@ def test_new_entities_pass_through_analyze_payload(): """ import json - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( mock_testing=True, pii_entities_config={ PiiEntityType.DE_TAX_ID: PiiAction.MASK, @@ -3634,7 +3634,7 @@ def _make_marker_session_iterator( def _chunking_guardrail(chunk_size_bytes=100, **kwargs): - return _OPTIONAL_PresidioPIIMasking( + return OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base="http://test-analyzer/", presidio_anonymizer_api_base="http://test-anonymizer/", presidio_analyze_chunk_size_bytes=chunk_size_bytes, @@ -3651,7 +3651,7 @@ def _oversized_marker_text(): def test_split_text_for_analysis_offsets_and_byte_budget(): text = " ".join(f"word{i}" for i in range(200)) - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20) + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20) assert len(chunks) > 1 for offset, chunk in chunks: assert len(chunk.encode("utf-8")) <= 100 @@ -3666,7 +3666,7 @@ def test_split_text_for_analysis_offsets_and_byte_budget(): def test_split_text_for_analysis_multibyte_characters(): text = "émoji🙂 çafé " * 120 - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=64, overlap_chars=8) + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=64, overlap_chars=8) assert len(chunks) > 1 for offset, chunk in chunks: assert len(chunk.encode("utf-8")) <= 64 @@ -3676,7 +3676,7 @@ def test_split_text_for_analysis_multibyte_characters(): def test_split_text_for_analysis_under_budget_returns_single_chunk(): text = "short text" - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20) + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20) assert chunks == [(0, text)] @@ -3805,18 +3805,18 @@ async def test_analyze_text_chunked_failure_stays_fail_closed(): def test_presidio_analyze_chunk_size_default_and_validation(): from litellm.constants import DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True) assert guardrail.presidio_analyze_chunk_size_bytes == DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES - nonpositive = _OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=-5) + nonpositive = OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=-5) assert nonpositive.presidio_analyze_chunk_size_bytes == DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES - custom = _OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=1234) + custom = OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=1234) assert custom.presidio_analyze_chunk_size_bytes == 1234 def test_update_in_memory_applies_analyze_chunk_size(): - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True) params = LitellmParams( guardrail="presidio", mode="pre_call", @@ -3827,8 +3827,8 @@ def test_update_in_memory_applies_analyze_chunk_size(): def test_update_in_memory_keeps_output_masker_from_unmasking(): - masker = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True, output_parse_pii=False) - unmasker = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) + masker = OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True, output_parse_pii=False) + unmasker = OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) params = LitellmParams(guardrail="presidio", mode="pre_call", output_parse_pii=True) masker.update_in_memory_litellm_params(params) @@ -3844,7 +3844,7 @@ def test_merge_drops_truncated_same_type_fragment_from_overlap(): numbered-token rewriter and double-counts entities.""" truncated = {"entity_type": "IP_ADDRESS", "start": 10, "end": 21, "score": 0.6} full_local = {"entity_type": "IP_ADDRESS", "start": 5, "end": 18, "score": 0.95} - merged = _OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( + merged = OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( text_chunks=[(0, "x" * 21), (5, "x" * 25)], chunk_results=[[truncated], [full_local]], ) @@ -3856,7 +3856,7 @@ def test_merge_drops_truncated_same_type_fragment_from_overlap(): def test_merge_exact_duplicate_keeps_higher_score(): low = {"entity_type": "EMAIL_ADDRESS", "start": 3, "end": 9, "score": 0.4} high = {"entity_type": "EMAIL_ADDRESS", "start": 0, "end": 6, "score": 0.9} - merged = _OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( + merged = OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( text_chunks=[(0, "x" * 9), (3, "x" * 9)], chunk_results=[[low], [high]], ) @@ -3869,7 +3869,7 @@ def test_merge_preserves_cross_type_overlap(): (e.g. URL inside EMAIL_ADDRESS); the chunk merge must not drop those.""" email = {"entity_type": "EMAIL_ADDRESS", "start": 0, "end": 20, "score": 1.0} url = {"entity_type": "URL", "start": 5, "end": 20, "score": 0.5} - merged = _OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( + merged = OPTIONAL_PresidioPIIMasking._merge_chunked_analyze_results( text_chunks=[(0, "x" * 25)], chunk_results=[[email, url]], ) @@ -3879,7 +3879,7 @@ def test_merge_preserves_cross_type_overlap(): def test_update_in_memory_coerces_invalid_chunk_size(): from litellm.constants import DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=99_000) + guardrail = OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=99_000) params = LitellmParams( guardrail="presidio", mode="pre_call", @@ -3890,7 +3890,7 @@ def test_update_in_memory_coerces_invalid_chunk_size(): def test_split_text_handles_chunk_size_below_char_width(): - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis( + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis( text="\U0001f642\U0001f642", chunk_size_bytes=3, overlap_chars=8 ) assert all(chunk for _, chunk in chunks) @@ -3967,7 +3967,7 @@ def test_split_text_accounts_for_json_body_expansion(): text = "これは個人情報テストです。" * 200 # 3-byte UTF-8 chars, 6-byte escapes budget = 1000 - chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=budget, overlap_chars=8) + chunks = OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=budget, overlap_chars=8) assert len(chunks) > 1 for offset, chunk in chunks: assert len(json_module.dumps(chunk).encode("utf-8")) - 2 <= budget @@ -4158,7 +4158,7 @@ async def test_pii_masking_replays_a_byte_identical_prefix_across_turns(mock_use the history lose their binding. The analyzer and anonymizer are an in-process fake handed to the guardrail through its api_base settings.""" async with TestServer(_fake_presidio_app()) as server: - guardrail = _OPTIONAL_PresidioPIIMasking( + guardrail = OPTIONAL_PresidioPIIMasking( presidio_analyzer_api_base=str(server.make_url("/")), presidio_anonymizer_api_base=str(server.make_url("/")), pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK}, @@ -4279,7 +4279,7 @@ async def test_restoration_never_contacts_presidio(has_tokens: bool) -> None: async def test_standalone_restoration_preserves_post_call_selection(event_hook: str | list[str]) -> None: from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails - callback: Final = _OPTIONAL_PresidioPIIMasking( + callback: Final = OPTIONAL_PresidioPIIMasking( event_hook=event_hook, default_on=True, output_parse_pii=True, @@ -4302,7 +4302,7 @@ async def test_standalone_restoration_preserves_post_call_selection(event_hook: ], ) def test_validate_environment_missing_http(base_url): - pii_masking = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + pii_masking = OPTIONAL_PresidioPIIMasking(mock_testing=True) env_vars = { "PRESIDIO_ANALYZER_API_BASE": f"{base_url}/analyze", @@ -4333,7 +4333,7 @@ async def test_output_parsing(): """ litellm.set_verbose = True litellm.output_parse_pii = True - pii_masking = _OPTIONAL_PresidioPIIMasking(mock_testing=True) + pii_masking = OPTIONAL_PresidioPIIMasking(mock_testing=True) initial_message = [ { @@ -4413,7 +4413,7 @@ async def test_presidio_pii_masking_input_a(): """ Tests to see if correct parts of sentence anonymized """ - pii_masking = _OPTIONAL_PresidioPIIMasking( + pii_masking = OPTIONAL_PresidioPIIMasking( mock_testing=True, mock_redacted_text=input_a_anonymizer_results ) @@ -4444,7 +4444,7 @@ async def test_presidio_pii_masking_input_b(): """ Tests to see if correct parts of sentence anonymized """ - pii_masking = _OPTIONAL_PresidioPIIMasking( + pii_masking = OPTIONAL_PresidioPIIMasking( mock_testing=True, mock_redacted_text=input_b_anonymizer_results ) @@ -4474,7 +4474,7 @@ async def test_presidio_pii_masking_input_b(): async def test_presidio_pii_masking_logging_output_only_no_pre_api_hook(): from litellm.types.guardrails import GuardrailEventHooks - pii_masking = _OPTIONAL_PresidioPIIMasking( + pii_masking = OPTIONAL_PresidioPIIMasking( logging_only=True, mock_testing=True, mock_redacted_text=input_b_anonymizer_results, @@ -4505,7 +4505,7 @@ async def test_presidio_language_configuration(): """Test that presidio_language parameter is properly set and used in analyze requests""" litellm.turn_on_debug() - presidio_guardrail_de = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail_de = OPTIONAL_PresidioPIIMasking( pii_entities_config={}, presidio_language="de", mock_testing=True, @@ -4520,7 +4520,7 @@ async def test_presidio_language_configuration(): assert analyze_request["language"] == "de" assert analyze_request["text"] == test_text - presidio_guardrail_es = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail_es = OPTIONAL_PresidioPIIMasking( pii_entities_config={}, presidio_language="es", mock_testing=True ) @@ -4533,7 +4533,7 @@ async def test_presidio_language_configuration(): assert analyze_request_es["language"] == "es" assert analyze_request_es["text"] == test_text_es - presidio_guardrail_default = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail_default = OPTIONAL_PresidioPIIMasking( pii_entities_config={}, mock_testing=True ) @@ -4554,7 +4554,7 @@ 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() - presidio_guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_guardrail = OPTIONAL_PresidioPIIMasking( pii_entities_config={}, presidio_language="de", mock_testing=True ) diff --git a/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py index 35137e4693e..17427e839b9 100644 --- a/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/unit/proxy/guardrails/test_deferred_guardrail_logging.py @@ -1114,7 +1114,7 @@ class TestDeferredStreamingClosure: with ( patch("litellm.callbacks", [guardrail]), patch( - "litellm.proxy.utils._check_and_merge_model_level_guardrails", + "litellm.proxy.utils.check_and_merge_model_level_guardrails", side_effect=mock_merge, ), ): @@ -1178,7 +1178,7 @@ class TestDeferredStreamingClosure: with ( patch("litellm.callbacks", [guardrail_a, guardrail_b]), patch( - "litellm.proxy.utils._check_and_merge_model_level_guardrails", + "litellm.proxy.utils.check_and_merge_model_level_guardrails", side_effect=mock_merge, ), ): @@ -1216,7 +1216,7 @@ class TestDeferredStreamingClosure: raise RuntimeError("Simulated init failure") with patch( - "litellm.proxy.utils._check_and_merge_model_level_guardrails", + "litellm.proxy.utils.check_and_merge_model_level_guardrails", side_effect=exploding_merge, ): await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( diff --git a/tests/unit/proxy/guardrails/test_guardrail_coverage.py b/tests/unit/proxy/guardrails/test_guardrail_coverage.py index 474ca9d8fa0..28e70ec84fa 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_coverage.py +++ b/tests/unit/proxy/guardrails/test_guardrail_coverage.py @@ -605,9 +605,9 @@ async def test_azure_content_safety_pre_call_fires_on_runtime_call_types( ``aresponses`` for the Responses API. The hook must inspect text fragments under both, not only the literal ``"completion"`` string used by some SDK callers.""" - from litellm.proxy.hooks.azure_content_safety import _PROXY_AzureContentSafety + from litellm.proxy.hooks.azure_content_safety import PROXY_AzureContentSafety - guard = _PROXY_AzureContentSafety.__new__(_PROXY_AzureContentSafety) + guard = PROXY_AzureContentSafety.__new__(PROXY_AzureContentSafety) seen = [] async def fake_test_violation(content, source=None): @@ -628,9 +628,9 @@ async def test_azure_content_safety_post_call_checks_all_choices(user_api_key): """Krrish blocker: ``n>1`` responses must not bypass Azure Content Safety by placing the unsafe text in ``choices[1+]``.""" from fastapi import HTTPException - from litellm.proxy.hooks.azure_content_safety import _PROXY_AzureContentSafety + from litellm.proxy.hooks.azure_content_safety import PROXY_AzureContentSafety - guard = _PROXY_AzureContentSafety.__new__(_PROXY_AzureContentSafety) + guard = PROXY_AzureContentSafety.__new__(PROXY_AzureContentSafety) seen_outputs = [] async def fake_test_violation(content, source=None): diff --git a/tests/unit/proxy/guardrails/test_init_guardrails.py b/tests/unit/proxy/guardrails/test_init_guardrails.py index 2cac5d60068..c8eeaf11fde 100644 --- a/tests/unit/proxy/guardrails/test_init_guardrails.py +++ b/tests/unit/proxy/guardrails/test_init_guardrails.py @@ -209,7 +209,7 @@ def test_initialize_presidio_forwards_analyze_chunk_size_bytes(): """ import litellm from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - _OPTIONAL_PresidioPIIMasking, + OPTIONAL_PresidioPIIMasking, ) test_guardrail = { @@ -229,7 +229,7 @@ def test_initialize_presidio_forwards_analyze_chunk_size_bytes(): initialized = [ callback for callback in litellm.callbacks - if isinstance(callback, _OPTIONAL_PresidioPIIMasking) and callback.guardrail_name == "test_presidio_chunk_size" + if isinstance(callback, OPTIONAL_PresidioPIIMasking) and callback.guardrail_name == "test_presidio_chunk_size" ] assert initialized, "presidio guardrail was not registered as a callback" assert initialized[-1].presidio_analyze_chunk_size_bytes == 250_000 diff --git a/tests/unit/proxy/health_endpoints/test_health_endpoints.py b/tests/unit/proxy/health_endpoints/test_health_endpoints.py index 95aad7b772f..266fd06333c 100644 --- a/tests/unit/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/unit/proxy/health_endpoints/test_health_endpoints.py @@ -426,7 +426,7 @@ async def test_test_model_connection_loads_config_from_router(): mock_run_with_timeout, ), patch( - "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + "litellm.proxy.health_endpoints._health_endpoints.update_litellm_params_for_health_check", mock_update_params, ), patch( @@ -575,7 +575,7 @@ async def test_test_model_connection_uses_model_info_id_to_disambiguate_duplicat mock_run_with_timeout, ), patch( - "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + "litellm.proxy.health_endpoints._health_endpoints.update_litellm_params_for_health_check", mock_update_params, ), patch( @@ -677,7 +677,7 @@ async def test_test_model_connection_falls_back_to_deployments_zero_without_id() mock_run_with_timeout, ), patch( - "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + "litellm.proxy.health_endpoints._health_endpoints.update_litellm_params_for_health_check", mock_update_params, ), patch( @@ -2566,13 +2566,13 @@ def test_no_federation_field_reaches_a_non_admin_health_entry(federation_field: deployment is healthy must learn neither. Both lists that enforce that are derived from the same key sets this runs over, so a field added to the funnel without joining either one shows up here as a value a non-admin could read.""" - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data from litellm.proxy.health_endpoints._health_endpoints import ( _strip_admin_only_fields_from_health_result, ) canary = f"CANARY-{federation_field}-VALUE" - cleaned = _clean_endpoint_data( + cleaned = clean_endpoint_data( {"model": "anthropic/claude-sonnet-5", federation_field: canary}, details=True, ) @@ -2591,10 +2591,10 @@ def test_no_federation_secret_reaches_even_an_admin_health_entry(secret_field: s token, key, or reference it federates with, so these fields drop at the health-check layer ahead of any per-caller stripping. Reading the same set the drop list is built from is what catches a new secret-bearing field that was only ever added to the admin-gated half.""" - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data canary = f"CANARY-{secret_field}-VALUE" - cleaned = _clean_endpoint_data( + cleaned = clean_endpoint_data( {"model": "anthropic/claude-sonnet-5", secret_field: canary}, details=True, ) @@ -3364,7 +3364,7 @@ def test_clean_endpoint_data_strips_credentials_keeps_routing_fields(): layer based on user role, not in the cleaning helper. This guarantees proxy admins continue to see those fields in the /health response. """ - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data raw = { "model": "openai/gpt-4o", @@ -3374,7 +3374,7 @@ def test_clean_endpoint_data_strips_credentials_keeps_routing_fields(): "aws_access_key_id": "AKIAEXAMPLE", } - cleaned = _clean_endpoint_data(raw, details=True) + cleaned = clean_endpoint_data(raw, details=True) assert "api_key" not in cleaned assert "aws_access_key_id" not in cleaned @@ -3388,7 +3388,7 @@ def test_clean_endpoint_data_strips_extra_headers_and_aws_session_token(): `extra_headers` / `headers` / `aws_session_token`. Before the fix these were returned in plaintext (api_key was stripped, but these were not). """ - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data raw = { "model": "openai/gpt-4o", @@ -3402,7 +3402,7 @@ def test_clean_endpoint_data_strips_extra_headers_and_aws_session_token(): "aws_session_token": "CANARY_AWS_SESSION_TOKEN_VALUE", } - cleaned = _clean_endpoint_data(raw, details=True) + cleaned = clean_endpoint_data(raw, details=True) assert "extra_headers" not in cleaned assert "headers" not in cleaned @@ -3439,10 +3439,10 @@ def test_clean_endpoint_data_never_displays_credential_fields(credential_field, LIT-6239 / gh-36898: /health entries, healthy and unhealthy alike, must never carry credential-bearing litellm_params, with or without details. """ - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data canary = f"CANARY-{credential_field}-VALUE" - cleaned = _clean_endpoint_data( + cleaned = clean_endpoint_data( { "model": "azure/gpt-5-mini", "api_base": "https://example.test/v1", @@ -4063,9 +4063,9 @@ def test_clean_endpoint_data_keeps_only_json_safe_diagnostics(): """ from fastapi.encoders import jsonable_encoder - from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_check import clean_endpoint_data - cleaned = _clean_endpoint_data( + cleaned = clean_endpoint_data( { "model": "bedrock/us.amazon.nova-2-lite-v1:0", "custom_llm_provider": "bedrock", diff --git a/tests/unit/proxy/hooks/test_batch_file_validation.py b/tests/unit/proxy/hooks/test_batch_file_validation.py index 38ee6997899..8eeee2e6468 100644 --- a/tests/unit/proxy/hooks/test_batch_file_validation.py +++ b/tests/unit/proxy/hooks/test_batch_file_validation.py @@ -151,9 +151,9 @@ async def test_pre_call_rejects_unauthorized_model_in_batch_file(): """Pre-fix the hook only validated the outer `model` parameter and forwarded the file as-is. With this fix, a model named inside the JSONL that the caller cannot use must trigger a 403.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -199,9 +199,9 @@ async def test_pre_call_allows_all_team_models_key_when_model_in_team_allowlist( """Keys with ``all-team-models`` must inherit the team allowlist when validating models embedded in batch JSONL.""" from litellm.proxy._types import SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -233,9 +233,9 @@ async def test_pre_call_allows_all_team_models_key_when_model_in_team_allowlist( @pytest.mark.asyncio async def test_pre_call_uses_current_team_allowlist_for_all_team_models_key(): from litellm.proxy._types import LiteLLM_TeamTable, SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -287,9 +287,9 @@ async def test_pre_call_allows_all_team_models_key_via_current_team_object(): allowlist must be authorized through the freshly-fetched team object, not the cached-``team_models`` fallback.""" from litellm.proxy._types import LiteLLM_TeamTable, SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -352,9 +352,9 @@ async def test_pre_call_denies_all_team_models_key_via_member_scope(): LiteLLM_TeamTable, SpecialModelNames, ) - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -416,9 +416,9 @@ async def test_pre_call_fails_closed_when_current_team_fetch_fails_for_all_team_ team_fetch_error, expected_status ): from litellm.proxy._types import SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -470,9 +470,9 @@ async def test_pre_call_allows_teamless_all_team_models_key(): someone re-introduces a teamless denial in _resolve_key_models_for_auth_check or adds a team_id guard that blocks the batch path.""" from litellm.proxy._types import SpecialModelNames - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -503,9 +503,9 @@ async def test_pre_call_allows_teamless_all_team_models_key(): async def test_pre_call_allows_authorized_model_in_batch_file(): """If every model in the JSONL is on the caller's allowlist, the hook must not raise.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -542,9 +542,9 @@ async def test_pre_call_allows_authorized_model_in_batch_file(): @pytest.mark.asyncio async def test_pre_call_skips_file_fetch_when_disabled_in_general_settings(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -562,14 +562,14 @@ async def test_pre_call_skips_file_fetch_when_disabled_in_general_settings(): ) assert result == {"input_file_id": "file-abc123"} - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.assert_not_called() + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.assert_not_called() @pytest.mark.asyncio async def test_pre_call_skips_file_fetch_for_configured_provider(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -600,7 +600,7 @@ async def test_pre_call_skips_file_fetch_for_configured_provider(): # work — assert the skip happened rather than the hook's error-recovery # path (which also returns data unchanged). mock_afile_content.assert_not_awaited() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.assert_not_called() + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.assert_not_called() @pytest.mark.asyncio @@ -609,9 +609,9 @@ async def test_pre_call_does_not_skip_for_spoofed_provider(): user-supplied ``custom_llm_provider`` that is not backed by the routing deployment must not trigger a skip: the input file must still be fetched and the rate-limit counters incremented.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -619,7 +619,7 @@ async def test_pre_call_does_not_skip_for_spoofed_provider(): # only thing that could prevent the fetch below is the provider skip. If the # spoofed ``custom_llm_provider`` were honored, afile_content would never be # awaited. - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 100}} ] rate_limiter.parallel_request_limiter.atomic_check_and_increment_by_n = AsyncMock( @@ -672,7 +672,7 @@ async def test_pre_call_does_not_skip_for_spoofed_provider(): async def test_count_input_file_usage_decodes_model_embedded_file_id(): import base64 - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter original_file_id = "file-provider-xyz" encoded_payload = ( @@ -684,7 +684,7 @@ async def test_count_input_file_usage_decodes_model_embedded_file_id(): ) encoded_file_id = f"file-{encoded_payload}" - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -727,9 +727,9 @@ async def test_pre_call_allows_stripped_provider_model_when_key_has_proxy_alias( """After replace_model_in_jsonl, body.model is the provider id (e.g. gpt-5.5). Auth must check target_model_names from the unified file id, not reverse-map the stripped id.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -791,9 +791,9 @@ async def test_pre_call_uses_target_model_names_not_stripped_reverse_lookup( ): """LIT-3593: three deployments strip to gpt-5.5; auth must use the upload target alias from target_model_names, not first-match reverse lookup.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -862,9 +862,9 @@ async def test_pre_call_uses_target_model_names_not_stripped_reverse_lookup( async def test_pre_call_skips_check_when_no_models_present(): """Files without any `body.model` (corrupt or empty) must not 500; the rate limiter logs a warning elsewhere and proceeds.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -889,9 +889,9 @@ async def test_pre_call_skips_check_when_no_models_present(): def _make_rate_limiter(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - return _PROXY_BatchRateLimiter( + return PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -954,7 +954,7 @@ def test_get_batch_routing_model_uses_unified_file_id_target(): return_value=None, ), patch( - "litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id", + "litellm.proxy.openai_files_endpoints.common_utils.is_base64_encoded_unified_file_id", return_value="unified-id", ), patch( @@ -969,9 +969,9 @@ def test_get_batch_routing_model_uses_unified_file_id_target(): def test_key_requires_batch_model_access_check_branches(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - check = _PROXY_BatchRateLimiter._key_requires_batch_model_access_check + check = PROXY_BatchRateLimiter._key_requires_batch_model_access_check assert check(UserAPIKeyAuth(api_key="sk", models=["*"])) is False assert check(UserAPIKeyAuth(api_key="sk", models=["all-proxy-models"])) is False assert ( @@ -1007,9 +1007,9 @@ def test_key_requires_batch_model_access_check_branches(): def test_has_applicable_batch_rate_limits(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - has_limits = _PROXY_BatchRateLimiter._has_applicable_batch_rate_limits + has_limits = PROXY_BatchRateLimiter._has_applicable_batch_rate_limits assert has_limits([{"rate_limit": {"tokens_per_unit": 100}}]) is True assert has_limits([{"rate_limit": {"requests_per_unit": 5}}]) is True assert has_limits([{"rate_limit": {"max_parallel_requests": 2}}]) is True @@ -1032,7 +1032,7 @@ def test_should_skip_ignores_client_supplied_metadata_flag(): body. The skip decision is server-controlled only, so with applicable rate limits the JSONL is still processed despite the client flag.""" rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1059,7 +1059,7 @@ def test_should_not_skip_for_forged_model_embedded_file_id(): import base64 rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1088,7 +1088,7 @@ def test_should_not_skip_for_skip_listed_top_level_model(): ``body.model`` entries. No per-model skip exists, so a skip-listed model over a plain file still gets processed.""" rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1114,7 +1114,7 @@ def test_should_not_skip_when_file_bound_provider_is_rate_limited(): import base64 rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1156,7 +1156,7 @@ def test_should_skip_when_file_bound_provider_is_skip_listed(): import base64 rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1194,7 +1194,7 @@ def test_warns_once_for_unsupported_model_skip_setting(): """Operators who set the no-op per-model skip key get a single warning so a misconfigured deployment does not silently leave batch limits unenforced.""" rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1221,7 +1221,7 @@ def test_warns_once_for_unsupported_model_skip_setting(): def test_no_warning_when_model_skip_setting_absent(): rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"requests_per_unit": 5}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1243,7 +1243,7 @@ def test_no_warning_when_model_skip_setting_absent(): def test_should_skip_when_no_rate_limits_configured(): rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {}} ] user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1261,7 +1261,7 @@ def test_should_skip_when_no_rate_limits_configured(): def test_should_not_skip_and_reuses_descriptors_when_limits_present(): rate_limiter = _make_rate_limiter() descriptors = [{"rate_limit": {"tokens_per_unit": 100}}] - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = ( + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = ( descriptors ) user = UserAPIKeyAuth(api_key="sk", models=["*"]) @@ -1357,17 +1357,17 @@ def test_resolve_fetch_params_model_embedded_fails_open_on_credential_error(): async def test_check_and_increment_computes_descriptors_when_not_passed(): from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) parallel_request_limiter = MagicMock() - parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {"tokens_per_unit": 100}} ] parallel_request_limiter.atomic_check_and_increment_by_n = AsyncMock( return_value={"overall_code": "OK", "statuses": []} ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=parallel_request_limiter, ) @@ -1379,7 +1379,7 @@ async def test_check_and_increment_computes_descriptors_when_not_passed(): descriptors=None, ) - parallel_request_limiter._create_rate_limit_descriptors.assert_called_once() + parallel_request_limiter.create_rate_limit_descriptors.assert_called_once() @pytest.mark.asyncio @@ -1390,17 +1390,17 @@ async def test_pre_call_enforces_project_otpm_limit_for_batch(): quota. The project OTPM descriptor must now be present and charged with the batch's estimated *output* tokens, not its input tokens.""" from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache() - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(local_cache), parallel_request_limiter=parallel_request_limiter, ) @@ -1448,17 +1448,17 @@ async def test_pre_call_enforces_project_itpm_limit_for_batch(): """Companion to the OTPM regression above: a project's ITPM quota must also apply to batch submissions.""" from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache() - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(local_cache), parallel_request_limiter=parallel_request_limiter, ) @@ -1505,17 +1505,17 @@ async def test_pre_call_enforces_project_otpm_limit_for_non_routing_row_model(): different, quota-limited model. That row's tokens must still be charged against its own model's project OTPM quota.""" from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache() - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(local_cache), parallel_request_limiter=parallel_request_limiter, ) @@ -1566,18 +1566,18 @@ async def test_pre_call_charges_each_row_model_against_its_own_project_quota(): model's request must succeed even though the over-limit model's row would fail on its own.""" from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PROJECT_OTPM_DESCRIPTOR_KEY, - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache() - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(local_cache), parallel_request_limiter=parallel_request_limiter, ) @@ -1650,7 +1650,7 @@ def test_should_not_skip_when_project_has_io_limit_for_non_routing_model(): rate_limiter = _make_rate_limiter() # No key/team/model-level limits at all -- only a project OTPM limit for a # model unrelated to the routing model below. - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {}} ] user = UserAPIKeyAuth( @@ -1673,7 +1673,7 @@ def test_should_skip_when_project_has_no_io_limits_and_no_other_limits(): with no ITPM/OTPM configuration anywhere must still get the fast-path skip when no other rate limits apply, exactly as before this fix.""" rate_limiter = _make_rate_limiter() - rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [ + rate_limiter.parallel_request_limiter.create_rate_limit_descriptors.return_value = [ {"rate_limit": {}} ] user = UserAPIKeyAuth( @@ -1693,9 +1693,9 @@ def test_should_skip_when_project_has_no_io_limits_and_no_other_limits(): @pytest.mark.asyncio async def test_count_input_file_usage_raises_on_non_bytes_content(): - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1800,9 +1800,9 @@ async def test_count_input_file_usage_streams_without_building_list(): """count_input_file_usage must count requests/tokens in one streaming pass. Mocks the download; asserts the count is correct and that the dict-list helper is never called (a revert to the list approach would call it).""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1852,9 +1852,9 @@ async def test_count_input_file_usage_enforces_models_when_token_counting_fails( NOT skip the model allowlist check. async_pre_call_hook swallows non-HTTP exceptions and submits the batch, so a raised counting error would otherwise fail open. The access check must still run and deny the restricted model.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1897,9 +1897,9 @@ async def test_count_input_file_usage_estimates_tokens_when_counting_fails_for_a zero the token total, which would let a caller evade the TPM limit by sending rows the counter cannot measure. The row falls back to a conservative size-based estimate so the batch proceeds with a non-zero count.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1941,9 +1941,9 @@ async def test_count_input_file_usage_collects_models_after_malformed_line(): named on a row AFTER a malformed line must still be collected and denied by the allowlist check, otherwise a caller could hide a restricted model behind a bad row.""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter - rate_limiter = _PROXY_BatchRateLimiter( + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=MagicMock(), ) @@ -1991,15 +1991,15 @@ def _output_estimator(): """A `_PROXY_BatchRateLimiter` whose output-token floor is observable: the no-`max_tokens` floor mock returns a distinctive sentinel so tests can tell "floor was used" apart from "an explicit cap was read".""" - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) limiter = MagicMock() limiter.no_max_tokens_output_floor.return_value = 999 - limiter.get_output_candidate_count = _PROXY_MaxParallelRequestsHandler_v3.get_output_candidate_count - return _PROXY_BatchRateLimiter( + limiter.get_output_candidate_count = PROXY_MaxParallelRequestsHandler_v3.get_output_candidate_count + return PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=limiter, ) @@ -2103,16 +2103,16 @@ def test_estimate_entry_output_tokens_multiplies_candidate_count(body_extra, exp def _enqueued_rate_limiter(): from litellm import DualCache - from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.batch_rate_limiter import PROXY_BatchRateLimiter from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache local_cache = DualCache(default_in_memory_ttl=60) internal_usage_cache = InternalUsageCache(local_cache) - parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache) - rate_limiter = _PROXY_BatchRateLimiter( + parallel_request_limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache) + rate_limiter = PROXY_BatchRateLimiter( internal_usage_cache=internal_usage_cache, parallel_request_limiter=parallel_request_limiter, ) diff --git a/tests/unit/proxy/hooks/test_batch_rate_limiter.py b/tests/unit/proxy/hooks/test_batch_rate_limiter.py index 930f62fcd10..a3e60a89c9f 100644 --- a/tests/unit/proxy/hooks/test_batch_rate_limiter.py +++ b/tests/unit/proxy/hooks/test_batch_rate_limiter.py @@ -19,7 +19,7 @@ from litellm.constants import BATCH_TPD_WINDOW_SECONDS from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.batch_rate_limiter import BatchFileUsage from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache, hash_token @@ -34,7 +34,7 @@ class _Clock: def _make_limiters(clock: _Clock | None = None): internal_usage_cache = InternalUsageCache(dual_cache=DualCache()) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache, time_provider=clock) + rate_limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache, time_provider=clock) batch_limiter = rate_limiter._get_batch_rate_limiter() assert batch_limiter is not None return internal_usage_cache, rate_limiter, batch_limiter @@ -252,7 +252,7 @@ def test_tpd_only_key_is_not_skipped_as_having_no_limits(): def test_online_descriptors_ignore_tpd_limit(): _internal_usage_cache, rate_limiter, _batch_limiter = _make_limiters() api_key = hash_token("online-key") - descriptors = rate_limiter._create_rate_limit_descriptors( + descriptors = rate_limiter.create_rate_limit_descriptors( user_api_key_dict=UserAPIKeyAuth(api_key=api_key, rpm_limit=5, tpd_limit=100, team_id="t", team_tpd_limit=9), data={"model": "gpt-4o"}, rpm_limit_type=None, diff --git a/tests/unit/proxy/hooks/test_dynamic_rate_limiter.py b/tests/unit/proxy/hooks/test_dynamic_rate_limiter.py index 611924ed658..cbb3e5f981c 100644 --- a/tests/unit/proxy/hooks/test_dynamic_rate_limiter.py +++ b/tests/unit/proxy/hooks/test_dynamic_rate_limiter.py @@ -5,9 +5,9 @@ import asyncio, importlib, litellm, os, pytest from litellm.caching.caching import DualCache from litellm.proxy.hooks.dynamic_rate_limiter import( - _PROXY_DynamicRateLimitHandler as DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler as DynamicRateLimitHandler, DynamicRateLimiterCache, - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) from litellm.types.utils import HiddenParams, ModelResponse from litellm import DualCache as DualCache_dynamic_rate, Router @@ -46,7 +46,7 @@ async def test_minute_rollover_between_sadd_and_get_reads_empty_window(): @pytest.mark.asyncio async def test_handler_threads_time_fn_to_internal_cache(): - handler = _PROXY_DynamicRateLimitHandler( + handler = PROXY_DynamicRateLimitHandler( internal_usage_cache=DualCache(), time_fn=lambda: datetime(2024, 1, 1, 10, 30, 0, tzinfo=timezone.utc), ) @@ -66,7 +66,7 @@ async def test_success_hook_updates_existing_hidden_params_storage() -> None: } ] ) - handler: Final = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) + handler: Final = PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) handler.update_variables(llm_router=router) response: Final = ModelResponse() hidden_params: Final = HiddenParams(model_id=model_id) diff --git a/tests/unit/proxy/hooks/test_dynamic_rate_limiter_v3.py b/tests/unit/proxy/hooks/test_dynamic_rate_limiter_v3.py index 527449bbc48..969c1eed63c 100644 --- a/tests/unit/proxy/hooks/test_dynamic_rate_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_dynamic_rate_limiter_v3.py @@ -17,7 +17,7 @@ import litellm from litellm import DualCache, Router from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3 as DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandlerV3 as DynamicRateLimitHandler, ) diff --git a/tests/unit/proxy/hooks/test_max_budget_per_session_limiter.py b/tests/unit/proxy/hooks/test_max_budget_per_session_limiter.py index a1b3f313814..45cc7eb8922 100644 --- a/tests/unit/proxy/hooks/test_max_budget_per_session_limiter.py +++ b/tests/unit/proxy/hooks/test_max_budget_per_session_limiter.py @@ -18,7 +18,7 @@ from litellm.caching.caching import DualCache from litellm.caching.redis_cache import _redis_circuit_breaker_guard from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.max_budget_per_session_limiter import ( - _PROXY_MaxBudgetPerSessionHandler, + PROXY_MaxBudgetPerSessionHandler, ) from litellm.proxy.utils import InternalUsageCache from litellm.types.agents import AgentResponse @@ -39,7 +39,7 @@ async def test_budget_per_session_under_budget_passes(): Requests under budget should pass through without error. """ local_cache = DualCache() - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -70,7 +70,7 @@ async def test_budget_per_session_exceeds_budget(): pre-call check should raise 429. """ local_cache = DualCache() - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -107,7 +107,7 @@ async def test_budget_per_session_independent_sessions(): Exhausting session A does not affect session B. """ local_cache = DualCache() - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -151,7 +151,7 @@ async def test_no_agent_id_passes(): When no agent_id is set on the key, all requests pass through. """ local_cache = DualCache() - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -190,7 +190,7 @@ class _OpenBreakerRedis: @pytest.mark.asyncio async def test_an_open_circuit_breaker_reads_session_spend_locally_without_a_warning(caplog): cache = DualCache(redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double - handler = _PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(cache)) + handler = PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(cache)) caplog.clear() with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): diff --git a/tests/unit/proxy/hooks/test_max_iterations_limiter.py b/tests/unit/proxy/hooks/test_max_iterations_limiter.py index 20928ef46d5..5eb499042b5 100644 --- a/tests/unit/proxy/hooks/test_max_iterations_limiter.py +++ b/tests/unit/proxy/hooks/test_max_iterations_limiter.py @@ -13,7 +13,7 @@ from fastapi import HTTPException from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler +from litellm.proxy.hooks.max_iterations_limiter import PROXY_MaxIterationsHandler from litellm.proxy.utils import InternalUsageCache from litellm.types.agents import AgentResponse @@ -36,7 +36,7 @@ async def test_max_iterations_basic_enforcement(): - 4th request should raise 429 """ local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( + handler = PROXY_MaxIterationsHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -46,9 +46,7 @@ async def test_max_iterations_basic_enforcement(): mock_agent = _make_mock_agent(max_iterations=3) - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: mock_registry.get_agent_by_id.return_value = mock_agent # First 3 requests should succeed @@ -81,7 +79,7 @@ async def test_max_iterations_different_sessions_independent(): - Exhausting Session A does not affect Session B """ local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( + handler = PROXY_MaxIterationsHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -91,9 +89,7 @@ async def test_max_iterations_different_sessions_independent(): mock_agent = _make_mock_agent(max_iterations=2) - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: mock_registry.get_agent_by_id.return_value = mock_agent # Session A: 2 calls succeed @@ -140,7 +136,7 @@ async def test_max_iterations_no_agent_id_passes(): When no agent_id is set on the key, all requests pass through. """ local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( + handler = PROXY_MaxIterationsHandler( internal_usage_cache=InternalUsageCache(local_cache), ) user_api_key_dict = UserAPIKeyAuth( diff --git a/tests/unit/proxy/hooks/test_model_max_budget_limiter.py b/tests/unit/proxy/hooks/test_model_max_budget_limiter.py index ffb60fb4651..8de34e6391e 100644 --- a/tests/unit/proxy/hooks/test_model_max_budget_limiter.py +++ b/tests/unit/proxy/hooks/test_model_max_budget_limiter.py @@ -8,7 +8,7 @@ import pytest from litellm.caching.caching import DualCache from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, ) from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.utils import LiteLLMBatch, Usage @@ -54,23 +54,23 @@ def _event(call_type: str, response_cost: float) -> dict[str, object]: } -async def _poll(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter, batch: LiteLLMBatch, response_cost: float) -> None: +async def _poll(limiter: PROXY_VirtualKeyModelMaxBudgetLimiter, batch: LiteLLMBatch, response_cost: float) -> None: await limiter.async_log_success_event( _event("aretrieve_batch", response_cost), response_obj=batch, start_time=None, end_time=None ) -async def _chat(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter) -> None: +async def _chat(limiter: PROXY_VirtualKeyModelMaxBudgetLimiter) -> None: await limiter.async_log_success_event( _event("acompletion", CHAT_COST), response_obj=None, start_time=None, end_time=None ) -async def _spend(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter, spend_key: str) -> float: +async def _spend(limiter: PROXY_VirtualKeyModelMaxBudgetLimiter, spend_key: str) -> float: return await limiter.dual_cache.async_get_cache(key=spend_key) or 0.0 -def _local_spend(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter, spend_key: str) -> float: +def _local_spend(limiter: PROXY_VirtualKeyModelMaxBudgetLimiter, spend_key: str) -> float: return limiter.dual_cache.in_memory_cache.get_cache(key=spend_key) or 0.0 @@ -133,8 +133,8 @@ class _SharedRedisDouble: return [await self.async_increment(op["key"], op["increment_value"], ttl=op["ttl"]) for op in increment_list] -def _worker(redis: _SharedRedisDouble) -> _PROXY_VirtualKeyModelMaxBudgetLimiter: - return _PROXY_VirtualKeyModelMaxBudgetLimiter( +def _worker(redis: _SharedRedisDouble) -> PROXY_VirtualKeyModelMaxBudgetLimiter: + return PROXY_VirtualKeyModelMaxBudgetLimiter( dual_cache=DualCache(redis_cache=redis) # pyright: ignore[reportArgumentType] # duck-typed Redis double ) @@ -145,7 +145,7 @@ async def _drain_redis_pushes() -> None: @pytest.mark.asyncio async def test_polls_of_a_finished_batch_charge_each_per_model_budget_once(): - limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) first: Final = _batch("batch_first", "completed") await _poll(limiter, _batch("batch_first", "in_progress"), response_cost=0) @@ -158,7 +158,7 @@ async def test_polls_of_a_finished_batch_charge_each_per_model_budget_once(): @pytest.mark.asyncio async def test_a_second_batch_and_chat_requests_still_charge_the_budget(): - limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) await _poll(limiter, _batch("batch_first", "completed"), response_cost=BATCH_COST) await _poll(limiter, _batch("batch_first", "completed"), response_cost=BATCH_COST) diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter.py b/tests/unit/proxy/hooks/test_parallel_request_limiter.py index c42d3b51799..1772f41a9d3 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter.py @@ -9,7 +9,7 @@ import pytest from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage @@ -17,7 +17,7 @@ from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage @pytest.mark.asyncio async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_the_login_token(): - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) session = UserAPIKeyAuth( api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa", user_id="alice", @@ -62,7 +62,7 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ team_id = "litellm-team" end_user_id = "customer-1" - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + parallel_request_handler = PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(DualCache()) ) diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index fecedd8c498..d9cf8536912 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -26,6 +26,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PARALLEL_REQUEST_SLOT_TTL_SECONDS, ParallelSlotAcquisition, + PROXY_MaxParallelRequestsHandler_v3, RateLimitDescriptor, RateLimitedModel, RateLimitResponse, @@ -36,7 +37,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( get_request_stash, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token from litellm.types.caching import RedisPipelineIncrementOperation @@ -99,7 +100,7 @@ def test_api_key_descriptor_applies_budget_throttle( budget_throttle_pct=throttle_pct, ) - descriptors = handler._create_rate_limit_descriptors( + descriptors = handler.create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data={}, rpm_limit_type=None, @@ -2914,7 +2915,7 @@ class TestGetTotalTokensFromUsageCacheExclusion: def handler(self): """Create a handler instance for testing.""" local_cache = DualCache() - return _PROXY_MaxParallelRequestsHandler( + return PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache), ) @@ -3786,9 +3787,9 @@ async def test_failure_event_settles_project_itpm_otpm_at_recovered_partial_usag # ----------------------- Per-MCP-server rate limiting (v3) ----------------------- -def _make_mcp_handler() -> tuple[_PROXY_MaxParallelRequestsHandler, DualCache]: +def _make_mcp_handler() -> tuple[PROXY_MaxParallelRequestsHandler_v3, DualCache]: local_cache: Final = DualCache() - handler: Final = _PROXY_MaxParallelRequestsHandler( + handler: Final = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=InternalUsageCache(local_cache) ) return handler, local_cache @@ -4878,7 +4879,7 @@ async def test_per_tag_rate_limit_independent_counters_v3(monkeypatch): @pytest.mark.asyncio async def test_per_tag_descriptor_creation_v3(): """ - _create_rate_limit_descriptors emits a tag_per_key descriptor carrying the + create_rate_limit_descriptors emits a tag_per_key descriptor carrying the configured RPM limit only for request tags present in the configured map. """ _api_key = hash_token("sk-per-tag-desc") @@ -4890,7 +4891,7 @@ async def test_per_tag_descriptor_creation_v3(): internal_usage_cache=InternalUsageCache(DualCache()) ) - descriptors = handler._create_rate_limit_descriptors( + descriptors = handler.create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1", "cell-2"]}}, rpm_limit_type=None, @@ -4916,7 +4917,7 @@ async def test_per_tag_descriptor_absent_without_config_v3(): internal_usage_cache=InternalUsageCache(DualCache()) ) - descriptors = handler._create_rate_limit_descriptors( + descriptors = handler.create_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1"]}}, rpm_limit_type=None, @@ -6425,7 +6426,7 @@ async def test_conflicting_token_limits_cannot_bypass_tpm_reservation(): def _enqueued_test_handler() -> _PROXY_MaxParallelRequestsHandler: - return _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60))) + return PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60))) def _batch_response(batch_id: str, status: str): @@ -6745,18 +6746,18 @@ def _handler_with_redis( ): internal_usage_cache = InternalUsageCache(DualCache(redis_cache=redis)) # pyright: ignore[reportArgumentType] # duck-typed Redis double if fail_closed is None and force_hash_tag_grouping is None: - return _PROXY_MaxParallelRequestsHandler(internal_usage_cache=internal_usage_cache) + return PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache) if fail_closed is None: - return _PROXY_MaxParallelRequestsHandler( + return PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache, force_hash_tag_grouping_resolver=lambda: force_hash_tag_grouping, ) if force_hash_tag_grouping is None: - return _PROXY_MaxParallelRequestsHandler( + return PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache, fail_closed_resolver=lambda: fail_closed, ) - return _PROXY_MaxParallelRequestsHandler( + return PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache, fail_closed_resolver=lambda: fail_closed, force_hash_tag_grouping_resolver=lambda: force_hash_tag_grouping, @@ -6853,7 +6854,7 @@ async def _admit(handler, auth, data=None): async def _read_only_check(handler, auth): - descriptors = handler._create_rate_limit_descriptors( + descriptors = handler.create_rate_limit_descriptors( user_api_key_dict=auth, data={"model": "test-model"}, rpm_limit_type=None, @@ -7675,7 +7676,7 @@ async def test_managed_invocations_enforce_actor_and_target_rate_policies( cache: Final = DualCache() handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) monkeypatch.setattr(handler, "_get_agent_from_registry", lambda _: None) - descriptors: Final = handler._create_rate_limit_descriptors( + descriptors: Final = handler.create_rate_limit_descriptors( user_api_key_dict=auth, data={"model": "a2a/target", "litellm_session_id": "session"}, rpm_limit_type=None, diff --git a/tests/unit/proxy/hooks/test_prompt_injection_detection.py b/tests/unit/proxy/hooks/test_prompt_injection_detection.py index b82b1be3b2a..3c871ca7333 100644 --- a/tests/unit/proxy/hooks/test_prompt_injection_detection.py +++ b/tests/unit/proxy/hooks/test_prompt_injection_detection.py @@ -11,7 +11,7 @@ import litellm from litellm.caching.caching import DualCache from litellm.proxy._types import LiteLLMPromptInjectionParams, UserAPIKeyAuth from litellm.proxy.hooks.prompt_injection_detection import ( - _OPTIONAL_PromptInjectionDetection, + OPTIONAL_PromptInjectionDetection, ) from litellm.proxy.utils import ProxyLogging from litellm.router import Router @@ -21,8 +21,8 @@ from litellm.utils import _invalidate_model_cost_lowercase_map from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome -def _moderation_detector(verdict: str) -> _OPTIONAL_PromptInjectionDetection: - detector = _OPTIONAL_PromptInjectionDetection( +def _moderation_detector(verdict: str) -> OPTIONAL_PromptInjectionDetection: + detector = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams( heuristics_check=False, llm_api_check=True, @@ -48,7 +48,7 @@ LONG_SAFE_PROMPT = "Summarize the quarterly revenue report for the finance team. @pytest.mark.asyncio async def test_acompletion_call_type_rejects_prompt_injection(): - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection() + prompt_injection_detection = OPTIONAL_PromptInjectionDetection() user_key = UserAPIKeyAuth(api_key="sk-test") cache = DualCache() data = { @@ -74,7 +74,7 @@ async def test_acompletion_call_type_rejects_prompt_injection(): @pytest.mark.asyncio async def test_acompletion_call_type_allows_safe_prompt(): - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection() + prompt_injection_detection = OPTIONAL_PromptInjectionDetection() user_key = UserAPIKeyAuth(api_key="sk-test") cache = DualCache() data = { @@ -153,7 +153,7 @@ async def test_proxy_during_call_hook_runs_configured_llm_api_check(monkeypatch) @pytest.mark.asyncio async def test_heuristics_check_keeps_event_loop_responsive(): - detector = _OPTIONAL_PromptInjectionDetection( + detector = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams(heuristics_check=True) ) data = {"model": "test-model", "messages": [{"role": "user", "content": LONG_SAFE_PROMPT}]} @@ -182,7 +182,7 @@ async def test_heuristics_check_keeps_event_loop_responsive(): @pytest.mark.asyncio async def test_heuristics_check_does_not_occupy_default_executor(): - detector = _OPTIONAL_PromptInjectionDetection( + detector = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams(heuristics_check=True) ) data = {"model": "test-model", "messages": [{"role": "user", "content": LONG_SAFE_PROMPT}]} @@ -329,7 +329,7 @@ async def test_prompt_injection_attack_valid_attack(): """ Tests if prompt injection detection catches a valid attack """ - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection() + prompt_injection_detection = OPTIONAL_PromptInjectionDetection() _api_key = "sk-98765" user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) @@ -360,7 +360,7 @@ async def test_prompt_injection_attack_invalid_attack(): Tests if prompt injection detection passes an invalid attack, which contains just 1 word """ litellm.set_verbose = True - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection() + prompt_injection_detection = OPTIONAL_PromptInjectionDetection() _api_key = "sk-98765" user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) @@ -398,7 +398,7 @@ async def test_prompt_injection_llm_eval(): llm_api_system_prompt="Detect if a prompt is safe to run. Return 'UNSAFE' if not.", llm_api_fail_call_string="UNSAFE", ) - prompt_injection_detection = _OPTIONAL_PromptInjectionDetection( + prompt_injection_detection = OPTIONAL_PromptInjectionDetection( prompt_injection_params=_prompt_injection_params, ) diff --git a/tests/unit/proxy/hooks/test_proxy_hooks_init.py b/tests/unit/proxy/hooks/test_proxy_hooks_init.py index 7f07fa7966c..a6edd3db944 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/hooks/test_proxy_rate_limit_provider_field.py b/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py index 49bbd498cb9..16b5406bc21 100644 --- a/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py +++ b/tests/unit/proxy/hooks/test_proxy_rate_limit_provider_field.py @@ -44,21 +44,21 @@ from litellm.exceptions import RateLimitError from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) -from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler +from litellm.proxy.hooks.dynamic_rate_limiter import PROXY_DynamicRateLimitHandler from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3, + PROXY_DynamicRateLimitHandlerV3, ) from litellm.proxy.hooks.max_budget_per_session_limiter import ( - _PROXY_MaxBudgetPerSessionHandler, + PROXY_MaxBudgetPerSessionHandler, ) -from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler +from litellm.proxy.hooks.max_iterations_limiter import PROXY_MaxIterationsHandler from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import ( @@ -166,9 +166,7 @@ class TestResolveLLMProviderForRateLimit: "litellm.proxy.proxy_server.llm_router", None, ): - resolved_model, provider = resolve_llm_provider_for_rate_limit( - "anything" - ) + resolved_model, provider = resolve_llm_provider_for_rate_limit("anything") assert provider == PROXY_LLM_PROVIDER_FALLBACK assert resolved_model == "anything" @@ -265,9 +263,7 @@ class TestResolveLLMProviderForRateLimit: "litellm.proxy.proxy_server.llm_router", _FakeRouter(), ): - resolved_model, provider = resolve_llm_provider_for_rate_limit( - "not-an-alias" - ) + resolved_model, provider = resolve_llm_provider_for_rate_limit("not-an-alias") assert provider == PROXY_LLM_PROVIDER_FALLBACK assert resolved_model == "not-an-alias" @@ -309,9 +305,7 @@ async def test_parallel_request_limiter_v1_populates_provider_when_at_rpm_limit( Trip the per-key RPM cap and assert the raised exception carries ``model`` / ``llm_provider`` resolved from ``data["model"]``. """ - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-test", max_parallel_requests=10, @@ -350,9 +344,7 @@ async def test_parallel_request_limiter_v1_zero_limit_path_populates_provider(): ``raise_rate_limit_error`` path. That path receives ``requested_model`` via the call-site change and must pass it through. """ - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-zero", max_parallel_requests=0, @@ -378,9 +370,7 @@ async def test_parallel_request_limiter_v1_zero_limit_path_populates_provider(): @pytest.mark.asyncio async def test_parallel_request_limiter_v1_global_limit_populates_provider(): """global_max_parallel_requests path also threads the model through.""" - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth(api_key="sk-global") # Pre-fill the global counter so the next call exceeds it. @@ -414,9 +404,7 @@ async def test_parallel_request_limiter_v1_unknown_model_falls_back(): When ``data["model"]`` is unparseable, the resolver falls back to ``litellm_proxy`` — and crucially does *not* leak a secondary exception. """ - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-unknown", max_parallel_requests=10, @@ -450,9 +438,7 @@ async def test_parallel_request_limiter_v1_unknown_model_falls_back(): @pytest.mark.asyncio async def test_parallel_request_limiter_v1_missing_model_falls_back(): - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) user_api_key_dict = UserAPIKeyAuth( api_key="sk-rl-no-model", max_parallel_requests=10, @@ -509,9 +495,7 @@ def _v3_over_limit_response(rate_limit_type: str = "requests") -> dict: ], ) async def test_parallel_request_limiter_v3_populates_provider(model, expected_provider): - handler = _PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}] over = _v3_over_limit_response() @@ -535,9 +519,7 @@ async def test_parallel_request_limiter_v3_populates_provider(model, expected_pr @pytest.mark.asyncio async def test_parallel_request_limiter_v3_unknown_model_falls_back(): - handler = _PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}] with pytest.raises(HTTPException) as exc_info: @@ -553,9 +535,7 @@ async def test_parallel_request_limiter_v3_unknown_model_falls_back(): @pytest.mark.asyncio async def test_parallel_request_limiter_v3_missing_model_falls_back(): - handler = _PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + handler = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}] with pytest.raises(HTTPException) as exc_info: @@ -576,7 +556,7 @@ async def test_parallel_request_limiter_v3_missing_model_falls_back(): @pytest.mark.asyncio async def test_dynamic_rate_limiter_v1_tpm_zero_populates_provider(): - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) handler.check_available_usage = AsyncMock(return_value=(0, 5, 100, 5, 1)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn") @@ -599,7 +579,7 @@ async def test_dynamic_rate_limiter_v1_tpm_zero_populates_provider(): @pytest.mark.asyncio async def test_dynamic_rate_limiter_v1_rpm_zero_populates_provider(): - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) handler.check_available_usage = AsyncMock(return_value=(5, 0, 5, 100, 1)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn") @@ -620,7 +600,7 @@ async def test_dynamic_rate_limiter_v1_rpm_zero_populates_provider(): @pytest.mark.asyncio async def test_dynamic_rate_limiter_v1_unknown_model_falls_back(): - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache()) handler.check_available_usage = AsyncMock(return_value=(0, 5, 100, 5, 1)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn") @@ -655,7 +635,7 @@ async def test_dynamic_rate_limiter_v3_model_capacity_path_populates_provider(): """ from litellm.types.router import ModelGroupInfo - handler = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache()) handler.v3_limiter.atomic_check_and_increment_by_n = AsyncMock( return_value={ "overall_code": "OVER_LIMIT", @@ -704,7 +684,7 @@ async def test_dynamic_rate_limiter_v3_unknown_descriptor_path_populates_provide """Fail-closed unknown-descriptor branch must still attribute provider.""" from litellm.types.router import ModelGroupInfo - handler = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache()) + handler = PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache()) handler.v3_limiter.atomic_check_and_increment_by_n = AsyncMock( return_value={ "overall_code": "OVER_LIMIT", @@ -773,16 +753,12 @@ async def test_batch_rate_limiter_populates_provider(): """ parallel_limiter = MagicMock() parallel_limiter.window_size = 60 - parallel_limiter._create_rate_limit_descriptors = MagicMock( - return_value=[ - {"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}} - ] - ) - parallel_limiter.atomic_check_and_increment_by_n = AsyncMock( - return_value=_batch_over_limit_response() + parallel_limiter.create_rate_limit_descriptors = MagicMock( + return_value=[{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}}] ) + parallel_limiter.atomic_check_and_increment_by_n = AsyncMock(return_value=_batch_over_limit_response()) - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(DualCache()), parallel_request_limiter=parallel_limiter, ) @@ -805,16 +781,12 @@ async def test_batch_rate_limiter_populates_provider(): async def test_batch_rate_limiter_unknown_model_falls_back(): parallel_limiter = MagicMock() parallel_limiter.window_size = 60 - parallel_limiter._create_rate_limit_descriptors = MagicMock( - return_value=[ - {"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}} - ] - ) - parallel_limiter.atomic_check_and_increment_by_n = AsyncMock( - return_value=_batch_over_limit_response() + parallel_limiter.create_rate_limit_descriptors = MagicMock( + return_value=[{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}}] ) + parallel_limiter.atomic_check_and_increment_by_n = AsyncMock(return_value=_batch_over_limit_response()) - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=InternalUsageCache(DualCache()), parallel_request_limiter=parallel_limiter, ) @@ -846,14 +818,10 @@ def _make_iter_agent(max_iterations: int) -> AgentResponse: @pytest.mark.asyncio async def test_max_iterations_limiter_populates_provider(): local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = PROXY_MaxIterationsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-iter", agent_id="agent-iter") - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1) await handler.async_pre_call_hook( @@ -887,14 +855,10 @@ async def test_max_iterations_limiter_populates_provider(): @pytest.mark.asyncio async def test_max_iterations_limiter_unknown_model_falls_back(): local_cache = DualCache() - handler = _PROXY_MaxIterationsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + handler = PROXY_MaxIterationsHandler(internal_usage_cache=InternalUsageCache(local_cache)) user_api_key_dict = UserAPIKeyAuth(api_key="sk-iter", agent_id="agent-iter") - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1) await handler.async_pre_call_hook( @@ -937,22 +901,12 @@ def _make_session_budget_agent(max_budget: float) -> AgentResponse: @pytest.mark.asyncio async def test_max_budget_per_session_limiter_populates_provider(): - handler = _PROXY_MaxBudgetPerSessionHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) - user_api_key_dict = UserAPIKeyAuth( - api_key="sk-session-budget", agent_id="agent-session-budget" - ) + handler = PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(DualCache())) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-session-budget", agent_id="agent-session-budget") - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: - mock_registry.get_agent_by_id.return_value = _make_session_budget_agent( - max_budget=1.0 - ) - with patch.object( - handler, "_get_current_spend", new=AsyncMock(return_value=5.0) - ): + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: + mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(max_budget=1.0) + with patch.object(handler, "_get_current_spend", new=AsyncMock(return_value=5.0)): with pytest.raises(HTTPException) as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -972,22 +926,12 @@ async def test_max_budget_per_session_limiter_populates_provider(): @pytest.mark.asyncio async def test_max_budget_per_session_limiter_unknown_model_falls_back(): - handler = _PROXY_MaxBudgetPerSessionHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) - user_api_key_dict = UserAPIKeyAuth( - api_key="sk-session-budget", agent_id="agent-session-budget" - ) + handler = PROXY_MaxBudgetPerSessionHandler(internal_usage_cache=InternalUsageCache(DualCache())) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-session-budget", agent_id="agent-session-budget") - with patch( - "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" - ) as mock_registry: - mock_registry.get_agent_by_id.return_value = _make_session_budget_agent( - max_budget=1.0 - ) - with patch.object( - handler, "_get_current_spend", new=AsyncMock(return_value=5.0) - ): + with patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry") as mock_registry: + mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(max_budget=1.0) + with patch.object(handler, "_get_current_spend", new=AsyncMock(return_value=5.0)): with pytest.raises(HTTPException) as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, @@ -1056,10 +1000,7 @@ def test_prometheus_exception_class_name_back_compat_for_budget_exceeded_error() # Default (empty llm_provider) path — same literal label. err_no_provider = litellm.BudgetExceededError(current_cost=1.0, max_budget=0.5) - assert ( - PrometheusLogger._get_exception_class_name(err_no_provider) - == "BudgetExceededError" - ) + assert PrometheusLogger._get_exception_class_name(err_no_provider) == "BudgetExceededError" if __name__ == "__main__": diff --git a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py index 7afa275c801..811176c4f4e 100644 --- a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py @@ -17,7 +17,7 @@ from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.db.spend_log_tool_index import response_tool_call_names from litellm.proxy.hooks.proxy_track_cost_callback import ( _get_budget_reservation_from_metadata, - _ProxyDBLogger, + ProxyDBLogger, _should_track_cost_callback, _update_database_and_spend_counters, run_spend_event, @@ -33,7 +33,7 @@ from litellm.types.utils import CallTypes, LiteLLMBatch, ModelResponse, Usage @pytest.mark.asyncio async def test_async_post_call_failure_hook(): # Setup - logger = _ProxyDBLogger() + logger = ProxyDBLogger() # Mock user_api_key_dict user_api_key_dict = UserAPIKeyAuth( @@ -103,7 +103,7 @@ async def test_async_post_call_failure_hook_carries_guardrail_info_from_litellm_ consume provider usage units, so the info must be carried over or the failure row logs guardrail_information: null. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() guardrail_info = [ { "guardrail_name": "bedrock-guard", @@ -136,7 +136,7 @@ async def test_async_post_call_failure_hook_carries_guardrail_info_from_litellm_ @pytest.mark.asyncio async def test_async_post_call_failure_hook_does_not_clobber_guardrail_info_in_metadata(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() metadata_bucket_info = [{"guardrail_name": "from-metadata-bucket"}] request_data = { "model": "gpt-4", @@ -173,7 +173,7 @@ async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from and leave request_data["metadata"] to the caller's native metadata, so a failed request on those routes wrote a spend row whose used_client_oauth_token was null instead of the stamped value """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "claude-sonnet-5", "custom_llm_provider": custom_llm_provider, @@ -217,7 +217,7 @@ async def test_async_post_call_failure_hook_never_lets_caller_metadata_set_used_ On /v1/messages and /v1/responses the request's own metadata field belongs to the caller, so a used_client_oauth_token they put there must never outrank the proxy's stamp or stand in for a missing one """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "claude-sonnet-5", "custom_llm_provider": "anthropic", @@ -264,7 +264,7 @@ async def test_async_post_call_failure_hook_never_lets_caller_metadata_set_used_ async def test_async_post_call_failure_hook_reads_used_client_oauth_token_from_the_routes_stamped_bucket( request_route: str, metadata_buckets: dict, expected: bool | None ): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "claude-sonnet-5", "custom_llm_provider": "anthropic", @@ -297,7 +297,7 @@ async def test_async_post_call_failure_hook_bills_guardrail_cost_on_blocked_requ """LIT-5651: a request blocked by a guardrail never reaches the LLM, but the guardrail invocation itself is billed by the provider. The failure row must charge that cost against the key instead of recording zero spend.""" - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], @@ -329,7 +329,7 @@ async def test_async_post_call_failure_hook_bills_guardrail_cost_on_blocked_requ @pytest.mark.asyncio async def test_async_post_call_failure_hook_adds_guardrail_cost_to_recovered_stream_cost(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() request_data = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], @@ -359,7 +359,7 @@ async def test_async_post_call_failure_hook_adds_guardrail_cost_to_recovered_str @pytest.mark.asyncio async def test_async_post_call_failure_hook_non_llm_route(): # Setup - logger = _ProxyDBLogger() + logger = ProxyDBLogger() # Mock user_api_key_dict with a non-LLM route user_api_key_dict = UserAPIKeyAuth( @@ -403,7 +403,7 @@ async def test_async_post_call_failure_hook_non_llm_route(): @pytest.mark.asyncio async def test_async_post_call_failure_hook_releases_budget_reservation_before_route_skip(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = {"reserved_cost": 0.5, "entries": []} user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -437,7 +437,7 @@ async def test_async_post_call_failure_hook_releases_budget_reservation_before_r @pytest.mark.asyncio async def test_should_continue_failure_tracking_when_budget_release_fails(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = {"reserved_cost": 0.5, "entries": []} user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -491,7 +491,7 @@ async def test_should_continue_failure_tracking_when_budget_release_fails(): @pytest.mark.asyncio async def test_track_cost_callback_releases_budget_reservation_when_spend_tracking_skips(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = {"reserved_cost": 0.5, "entries": []} user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) @@ -527,7 +527,7 @@ async def test_track_cost_callback_releases_budget_reservation_when_spend_tracki @pytest.mark.asyncio async def test_track_cost_callback_releases_budget_reservation_when_response_cost_missing(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = {"reserved_cost": 0.5, "entries": []} user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) @@ -970,7 +970,7 @@ async def test_track_cost_callback_skips_when_no_standard_logging_object(): File operations have no model and no standard_logging_object. The callback should skip gracefully instead of raising. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": "afile_delete", @@ -1011,7 +1011,7 @@ async def test_track_cost_callback_defers_in_progress_background_interaction(): """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": "acreate_interaction", @@ -1105,7 +1105,7 @@ async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # are gated, since creating a batch is its own billable request, and a retrieve that charges nothing hands its budget reservation back instead. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() budget_reservation = None if charged else {"reserved_cost": 0.5, "entries": []} kwargs = _batch_retrieve_kwargs(call_type, reservation=budget_reservation) @@ -1170,7 +1170,7 @@ async def test_track_cost_callback_keeps_reservation_open_for_in_progress_backgr """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} in_progress_response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1210,7 +1210,7 @@ async def test_track_cost_callback_releases_reservation_for_in_progress_interact monkeypatch.setattr(callback_module, "BACKGROUND_INTERACTION_COST_POLLING_ENABLED", False) - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} in_progress_response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1256,7 +1256,7 @@ async def test_track_cost_callback_releases_reservation_for_unpollable_interacti """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} terminal_response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1297,7 +1297,7 @@ async def test_track_cost_callback_alerts_when_an_interaction_that_produced_outp """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} usageless_response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1333,7 +1333,7 @@ async def test_track_cost_callback_releases_reservation_for_interaction_without_ """ from litellm.types.interactions import InteractionsAPIResponse - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} idless_response = InteractionsAPIResponse( id="", @@ -1379,7 +1379,7 @@ async def test_callback_handles_every_status_the_interactions_api_can_return(): released = set() for status in sorted(member.value for member in Status1): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() reservation = {"reserved_cost": 0.05, "entries": [], "finalized": False} response = InteractionsAPIResponse( id="interactions/bg-abc", @@ -1425,7 +1425,7 @@ async def test_async_post_call_failure_hook_propagates_trace_id_from_logging_obj The failure hook should propagate this so the DB spend log's session_id matches the Langfuse trace_id. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -1494,7 +1494,7 @@ async def test_enrich_failure_metadata_with_team_alias(): "user_api_key_team_id": "test_team_id", "user_api_key_team_alias": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) assert result["user_api_key_team_alias"] == "my-team-alias" @@ -1536,7 +1536,7 @@ async def test_enrich_failure_metadata_with_full_key_lookup(): "user_api_key_org_id": None, "user_api_key_project_id": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) assert result["user_api_key_alias"] == "fetched-key-alias" assert result["user_api_key_user_id"] == "fetched-user-id" assert result["user_api_key_team_id"] == "fetched-team-id" @@ -1567,7 +1567,7 @@ async def test_enrich_failure_metadata_skips_when_team_alias_present(): "user_api_key_team_id": "test_team_id", "user_api_key_team_alias": "already-set", } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) assert result["user_api_key_team_alias"] == "already-set" mock_get_key.assert_not_called() mock_get_team.assert_not_called() @@ -1589,7 +1589,7 @@ async def test_enrich_failure_metadata_skips_when_no_api_key(): "user_api_key_team_id": None, "user_api_key_team_alias": None, } - await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + await ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) mock_get_key.assert_not_called() @@ -1629,7 +1629,7 @@ async def test_enrich_failure_metadata_keeps_captured_identity_when_not_resolvin "user_api_key_team_alias": None, "user_api_key_org_id": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info( metadata, resolve_missing_key_identity=False ) @@ -1670,7 +1670,7 @@ async def test_enrich_failure_metadata_ignores_flag_when_alias_present(): "user_api_key_team_alias": None, "user_api_key_org_id": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + result = await ProxyDBLogger._enrich_failure_metadata_with_key_info( metadata, resolve_missing_key_identity=resolve ) mock_get_key.assert_not_called() @@ -1693,7 +1693,7 @@ async def test_track_cost_callback_reads_key_only_for_in_request_logs(call_type, identity persisted at create time. Every other call type still backfills from the key. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() mock_key_obj = MagicMock() mock_key_obj.key_alias = "alias-assigned-later" @@ -1760,7 +1760,7 @@ async def test_async_post_call_failure_hook_enriches_auth_error_metadata(): UserAPIKeyAuth is created with only api_key set. The failure hook should look up the key and team from cache/DB to populate all missing fields. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() # This is what auth_exception_handler creates for 401 errors user_api_key_dict = UserAPIKeyAuth( @@ -1817,7 +1817,7 @@ async def test_async_post_call_failure_hook_enriches_auth_error_metadata(): @pytest.mark.asyncio async def test_async_post_call_failure_hook_skips_the_key_lookup_when_the_failure_is_a_db_stall(): - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth(api_key="hashed_key") request_data = { "model": "gpt-5.6", @@ -1859,7 +1859,7 @@ async def test_async_post_call_failure_hook_skips_the_key_lookup_when_the_failur async def test_async_post_call_failure_hook_still_enriches_metadata_for_a_non_stall_failure(): """Only a DBLookupDeadlineExceeded skips the key lookup; a transport error from the provider call must still resolve the key's alias for the failure row.""" - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth(api_key="hashed_key") request_data = { "model": "gpt-5.6", @@ -1908,7 +1908,7 @@ async def test_async_post_call_failure_hook_enriches_missing_team_alias(): should look up the team from cache and populate user_api_key_team_alias in the spend log metadata written to the DB. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -1959,7 +1959,7 @@ async def test_track_cost_callback_skips_for_falsy_model_and_no_slo(model_value) Same bug as above but model can also be empty string (e.g. health check callbacks). The guard should catch all falsy model values when sl_object is missing. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": "acompletion", @@ -1996,7 +1996,7 @@ async def test_async_post_call_failure_hook_uses_actual_start_time(): """ from datetime import timedelta - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", @@ -2055,7 +2055,7 @@ async def _invoke_failure_hook_with_raised_exception(): Returns the metadata dict that was forwarded to ``update_database`` so the caller can assert on its ``error_information`` payload. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", user_id="u", @@ -2135,7 +2135,7 @@ async def test_async_post_call_failure_hook_records_recovered_partial_spend(): """ from litellm.types.utils import Usage - logger = _ProxyDBLogger() + logger = ProxyDBLogger() user_api_key_dict = UserAPIKeyAuth(api_key="test_api_key", user_id="u", team_id="t") request_data = { @@ -2166,7 +2166,7 @@ async def test_track_cost_callback_enriches_user_id_for_mcp_style_metadata(): """MCP tool calls may only carry user_api_key; user/team rollups still need user_id.""" from litellm.proxy._types import UserAPIKeyAuth - logger = _ProxyDBLogger() + logger = ProxyDBLogger() key_obj = UserAPIKeyAuth( api_key="hashed-key", user_id="mcp-user@example.com", @@ -2235,7 +2235,7 @@ async def test_track_cost_callback_keeps_guardrail_cost_on_cache_hit(): guardrail's provider charge must still reach spend logs and budgets. The payload already prices the LLM share at 0 on a cache hit, so its response_cost is the guardrail cost alone and the callback must pass it through untouched.""" - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": "acompletion", "model": "gpt-4o", @@ -2354,7 +2354,7 @@ async def test_track_cost_callback_logs_unauthenticated_pass_through_request(cal aretrieve_batch is included because CheckBatchCost's completed-batch cost event reaches this same callback with no attributable key/user/team when the batch was created with the master key or a team-less key.""" - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "call_type": call_type, @@ -2425,7 +2425,7 @@ async def _groups_charged_by_the_callback(kwargs, deployments=None): The callback resolves ``proxy_logging_obj`` and the router by importing them off ``proxy_server`` inside its own body, so there is no seam to inject either through. """ - logger = _ProxyDBLogger() + logger = ProxyDBLogger() with ( patch( # test-quality-ok: callback imports proxy_logging_obj off proxy_server in its body, no seam "litellm.proxy.proxy_server.proxy_logging_obj" @@ -2606,7 +2606,7 @@ async def test_async_log_success_event_hands_the_sidecar_a_compact_event_and_ski producer = SpendEventProducer( address=address, on_unavailable="fallback", buffer_size=10, connect_timeout=1.0, fallback=_no_fallback ) - logger = _ProxyDBLogger(producer) + logger = ProxyDBLogger(producer) with ( patch( # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam @@ -2642,7 +2642,7 @@ async def test_async_log_success_event_keeps_batch_retrieves_in_process(): connect_timeout=1.0, fallback=_no_fallback, ) - logger = _ProxyDBLogger(producer) + logger = ProxyDBLogger(producer) kwargs = {**_offload_kwargs(), "call_type": CallTypes.aretrieve_batch.value} completed_batch = LiteLLMBatch( id="batch_abc", @@ -2709,7 +2709,7 @@ async def test_sidecar_writes_the_same_spend_row_and_counters_as_the_in_process_ end_time = datetime(2026, 1, 1, 0, 0, 2) async def in_process() -> None: - await _ProxyDBLogger().async_log_success_event(_offload_kwargs(), _offload_response(), start_time, end_time) + await ProxyDBLogger().async_log_success_event(_offload_kwargs(), _offload_response(), start_time, end_time) async def via_sidecar() -> None: line = build_spend_event(_offload_kwargs(), _offload_response(), start_time, end_time, store_bodies=False) @@ -2753,7 +2753,7 @@ async def test_async_post_call_failure_hook_persists_no_raw_model_on_an_unknown_ raw_model: Final = "opus-4.6 Please summarize my medical records\nPatient has diabetes" writer: Final = MagicMock(spec=DBSpendUpdateWriter) writer.update_database = AsyncMock() - logger: Final = _ProxyDBLogger(spend_writer=lambda: writer) + logger: Final = ProxyDBLogger(spend_writer=lambda: writer) await logger.async_post_call_failure_hook( request_data={"model": raw_model, "messages": [{"role": "user", "content": "hi"}]}, @@ -2800,7 +2800,7 @@ def _spend_write_kwargs_with_metadata_value(metadata_value: object) -> dict: @pytest.mark.asyncio @pytest.mark.parametrize("log_level", [logging.WARNING, logging.DEBUG]) async def test_track_cost_callback_failure_alert_never_carries_request_metadata_values(log_level): - logger: Final = _ProxyDBLogger() + logger: Final = ProxyDBLogger() records: list[logging.LogRecord] = [] handler: Final = logging.Handler() handler.emit = records.append @@ -2862,7 +2862,7 @@ async def test_autonomous_llm_callback_persists_without_human_or_key(identity_fi new_callable=AsyncMock, return_value=False, ) as persist: - await _ProxyDBLogger()._PROXY_track_cost_callback( + await ProxyDBLogger()._PROXY_track_cost_callback( kwargs=kwargs, completion_response=ModelResponse(), start_time=datetime.now(), end_time=datetime.now() ) persist.assert_awaited_once() @@ -2889,7 +2889,7 @@ async def test_track_cost_callback_enqueue_emits_no_service_span(): # test-qual emitted; the flush that writes the queue emits its own table-named spans.""" from litellm.proxy.proxy_server import proxy_logging_obj - logger = _ProxyDBLogger() + logger = ProxyDBLogger() kwargs = { "model": "gpt-4", "call_type": "acompletion", diff --git a/tests/unit/proxy/hooks/test_rate_limiter_toctou.py b/tests/unit/proxy/hooks/test_rate_limiter_toctou.py index 23c717b0e3a..6a767a6486f 100644 --- a/tests/unit/proxy/hooks/test_rate_limiter_toctou.py +++ b/tests/unit/proxy/hooks/test_rate_limiter_toctou.py @@ -28,10 +28,10 @@ from litellm import DualCache, Router from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.batch_rate_limiter import BatchFileUsage from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3 as DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandlerV3 as DynamicRateLimitHandler, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache, hash_token @@ -89,7 +89,7 @@ async def test_batch_limiter_concurrent_bypasses_tpm_via_toctou(): dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() @@ -142,7 +142,7 @@ async def test_batch_limiter_uses_atomic_check_and_increment(): """ dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() @@ -356,7 +356,7 @@ async def test_batch_zero_token_consumes_rpm_only(): """ dual_cache = DualCache() internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + rate_limiter = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=internal_usage_cache ) batch_limiter = rate_limiter._get_batch_rate_limiter() diff --git a/tests/unit/proxy/hooks/test_sensitive_data_routing.py b/tests/unit/proxy/hooks/test_sensitive_data_routing.py index 48a42e42c62..fbe97e2917a 100644 --- a/tests/unit/proxy/hooks/test_sensitive_data_routing.py +++ b/tests/unit/proxy/hooks/test_sensitive_data_routing.py @@ -24,7 +24,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.sensitive_data_routing import ( DEFAULT_SENSITIVE_ROUTING_TTL, SENSITIVE_ROUTING_CACHE_PREFIX, - _PROXY_SensitiveDataRoutingHandler, + PROXY_SensitiveDataRoutingHandler, ) from litellm.proxy.utils import InternalUsageCache @@ -48,7 +48,7 @@ class TestSensitiveDataRoutingHandler: @pytest.fixture def handler(self): cache = MockInternalUsageCache() - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.fixture def user_api_key_dict(self): @@ -87,7 +87,7 @@ class TestSensitiveDataRoutingHandler: await super().async_set_cache(key, value, ttl=ttl, **kwargs) cache = TargetRecordingCache() - handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + handler = PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) await handler.set_session_routing( session_id="s-1", model="on-premise-model", user_api_key_dict=user_api_key_dict, guardrail_name="g" ) @@ -320,7 +320,7 @@ class TestStickySessionRouting: @pytest.fixture def handler(self): cache = MockInternalUsageCache() - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.fixture def user_api_key_dict(self): @@ -436,37 +436,37 @@ class TestCacheKeyAndTTL: def test_make_cache_key_format(self): cache = MockInternalUsageCache() - handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + handler = PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) key = handler._make_cache_key("test-session-123", "hashed-key") assert key == "{sensitive_route:hashed-key:test-session-123}:model" def test_make_cache_key_is_tenant_scoped(self): cache = MockInternalUsageCache() - handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + handler = PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) key_a = handler._make_cache_key("shared-session", "key-a") key_b = handler._make_cache_key("shared-session", "key-b") assert key_a != key_b def test_resolve_tenant_prefers_api_key(self): - tenant = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( + tenant = PROXY_SensitiveDataRoutingHandler._resolve_tenant( UserAPIKeyAuth(api_key="hashed-key", user_id="alice") ) assert tenant == "hashed-key" def test_resolve_tenant_falls_back_to_jwt_principal(self): - tenant = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( + tenant = PROXY_SensitiveDataRoutingHandler._resolve_tenant( UserAPIKeyAuth(api_key=None, user_id="alice", team_id="t1", org_id="o1") ) assert tenant == "user:alice|team:t1|org:o1" def test_resolve_tenant_distinguishes_keyless_principals(self): - tenant_a = _PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="alice")) - tenant_b = _PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="bob")) + tenant_a = PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="alice")) + tenant_b = PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="bob")) assert tenant_a != tenant_b def test_resolve_tenant_defaults_when_anonymous(self): - assert _PROXY_SensitiveDataRoutingHandler._resolve_tenant(None) == "default" - assert _PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None)) == "default" + assert PROXY_SensitiveDataRoutingHandler._resolve_tenant(None) == "default" + assert PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None)) == "default" class TestCustomGuardrailSessionIdExtraction: @@ -557,7 +557,7 @@ class TestRedisCache: cache = MockInternalUsageCache() mock_redis = AsyncMock() cache.dual_cache.redis_cache = mock_redis - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.mark.asyncio async def test_get_routed_model_from_redis(self, handler_with_redis): @@ -652,7 +652,7 @@ class TestPreCallHookEdgeCases: @pytest.fixture def handler(self): cache = MockInternalUsageCache() - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.fixture def user_api_key_dict(self): @@ -715,7 +715,7 @@ class TestProxyHandleSensitiveDataRouteException: @pytest.fixture def routing_hook(self): cache = MockInternalUsageCache() - return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + return PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) @pytest.mark.asyncio async def test_sticky_routing_persists_override(self, proxy_logging, routing_hook): @@ -1022,7 +1022,7 @@ class _OpenBreakerRedis: @pytest.mark.asyncio async def test_an_open_circuit_breaker_keeps_session_routing_in_memory_without_a_warning(caplog): cache = DualCache(redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double - handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=InternalUsageCache(cache)) + handler = PROXY_SensitiveDataRoutingHandler(internal_usage_cache=InternalUsageCache(cache)) caplog.clear() with caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"): diff --git a/tests/unit/proxy/hooks/test_tpm_concurrent.py b/tests/unit/proxy/hooks/test_tpm_concurrent.py index 42c1f489bdd..2d23b134433 100644 --- a/tests/unit/proxy/hooks/test_tpm_concurrent.py +++ b/tests/unit/proxy/hooks/test_tpm_concurrent.py @@ -27,7 +27,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PROJECT_OTPM_DESCRIPTOR_KEY, RateLimitedModel, _AUDIO_BYTES_PER_TOKEN, - _PROXY_MaxParallelRequestsHandler_v3 as RateLimitHandler, + PROXY_MaxParallelRequestsHandler_v3 as RateLimitHandler, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _call_id_from_callback_kwargs, diff --git a/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py b/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py index f194e43c74a..dc9dac9a57b 100644 --- a/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py +++ b/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py @@ -14,7 +14,7 @@ from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy._types import Litellm_EntityType from litellm.proxy.hooks.model_max_budget_limiter import ( _budget_model_candidates, - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, build_model_max_budget_usage, resolve_model_budget, ) @@ -27,7 +27,7 @@ from litellm.types.utils import BudgetConfig as GenericBudgetInfo @pytest.fixture def budget_limiter(): dual_cache = DualCache() - return _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + return PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) # Test _budget_model_candidates @@ -462,7 +462,7 @@ async def test_async_log_success_event_pushes_redis_increments_when_redis_config """ dual_cache = DualCache() dual_cache.redis_cache = object() # truthy placeholder; push only checks is not None - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) model = "gpt-4" kwargs = { "standard_logging_object": { @@ -491,7 +491,7 @@ async def test_async_log_success_event_pushes_redis_increments_when_redis_config @pytest.mark.asyncio async def test_model_budget_limiter_initializes_redis_increment_queue_lock(): dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) spend_key = "virtual_key_spend:test-key:gpt-4:1d" await limiter._increment_spend_in_current_window( @@ -653,7 +653,7 @@ async def test_logged_spend_is_visible_to_key_info_usage_and_enforcement(request actively blocked at 429 while reporting current_spend 0. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) key_hash = "vk-hash" model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} @@ -717,7 +717,7 @@ async def test_user_model_budget_is_tracked_and_enforced(): enforced, independently of any key-level budget. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) user_id = "user-1" user_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1mo"}} @@ -765,7 +765,7 @@ async def test_user_model_budget_counter_is_separate_from_the_key_counter(): counters, so one request must charge each exactly once. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.async_log_success_event( @@ -793,7 +793,7 @@ async def test_two_models_on_one_key_do_not_share_a_budget_window(): per model: a shared start lets the shorter period restart the longer one. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) model_max_budget = { "gpt-4": {"budget_limit": 10.0, "time_period": "1d"}, "claude-3": {"budget_limit": 10.0, "time_period": "30d"}, @@ -824,7 +824,7 @@ async def test_two_models_on_one_key_do_not_share_a_budget_window(): @pytest.mark.asyncio async def test_no_increment_when_no_scope_budgets_the_model(): dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) with patch.object(limiter, "_increment_spend_for_key", new_callable=AsyncMock) as mock_increment: await limiter.async_log_success_event( _success_kwargs( @@ -869,7 +869,7 @@ async def test_bedrock_traffic_charges_the_bare_family_name_budget(): the key went. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) key_hash = "vk-hash" model_max_budget = {"claude-opus-4-8": {"budget_limit": 1.0, "time_period": "18h"}} user_api_key = UserAPIKeyAuth(token=key_hash, model_max_budget=model_max_budget) @@ -916,7 +916,7 @@ async def test_user_model_budget_window_resets_when_the_period_elapses(): ) dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) user_id = "user-1" user_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1mo"}} spend_key = model_budget_spend_cache_key( @@ -976,7 +976,7 @@ async def test_a_zero_dollar_cap_blocks_the_model(): mean something. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) key = UserAPIKeyAuth( token="hash-zero", model_max_budget={"gpt-4": {"budget_limit": 0, "time_period": "1d"}}, @@ -1006,7 +1006,7 @@ async def test_spend_exactly_at_the_cap_is_refused(): (RouterBudgetLimiting, the key and team budget checks) uses `>=`. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) budget = {"gpt-4": {"budget_limit": 2.0, "time_period": "1d"}} key = UserAPIKeyAuth(token="hash-exact", model_max_budget=budget) @@ -1075,7 +1075,7 @@ async def test_one_malformed_scope_does_not_abort_the_other_scopes(): charged despite the user's entry being garbage. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) key_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.async_log_success_event( @@ -1107,7 +1107,7 @@ async def test_an_unusable_budget_entry_is_not_enforced_instead_of_raising(): cannot be keyed, so it cannot be enforced; the write path rejects these, so reaching here means config.yaml or a direct DB edit. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) key = UserAPIKeyAuth( token="hash-malformed", model_max_budget={"gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}}, @@ -1151,7 +1151,7 @@ def test_a_malformed_specific_entry_does_not_hide_a_usable_family_budget(): async def test_a_malformed_specific_entry_still_enforces_the_family_budget(): """The fall-through has to reach enforcement, not just resolution.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) budget = { "openai/gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}, "gpt-4": {"budget_limit": 1.0, "time_period": "1d"}, @@ -1228,7 +1228,7 @@ async def test_a_pre_upgrade_counter_keyed_on_the_request_model_still_enforces(e configured-model key finds that counter empty and admits another full budget until the window expires. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.dual_cache.async_set_cache(key=f"{prefix}:entity-1:openai/gpt-4:1d", value=25.0, ttl=86400) @@ -1264,7 +1264,7 @@ async def test_the_pre_upgrade_and_post_upgrade_counters_add_up_over_one_window( under-reports the window: 6 + 5 is over a cap of 10 that neither half reaches on its own. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.dual_cache.async_set_cache( key="virtual_key_spend:entity-1:openai/gpt-4:1d", value=legacy_spend, ttl=86400 @@ -1293,7 +1293,7 @@ async def test_the_configured_model_counter_is_never_counted_twice(): added them without noticing would charge 12 against a cap of 10 and refuse a key that has spent 6. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) await limiter.dual_cache.async_set_cache(key="virtual_key_spend:entity-1:gpt-4:1d", value=6.0, ttl=86400) assert ( @@ -1318,7 +1318,7 @@ async def test_the_pre_upgrade_counter_is_no_longer_read_a_window_after_start_up """ import litellm.proxy.hooks.model_max_budget_limiter as limiter_module - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.dual_cache.async_set_cache(key="virtual_key_spend:entity-1:openai/gpt-4:1d", value=25.0, ttl=86400) user_api_key = UserAPIKeyAuth(token="entity-1", model_max_budget=model_max_budget) @@ -1339,7 +1339,7 @@ async def test_the_user_scope_has_no_pre_upgrade_counter_to_carry(): Reading one would invent a counter no previous version ever wrote, which is the opposite of preserving one. """ - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) model_max_budget = {"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}} await limiter.dual_cache.async_set_cache(key="user_model_spend:u1:openai/gpt-4:1d", value=25.0, ttl=86400) @@ -1414,8 +1414,8 @@ async def test_spend_logged_on_one_replica_is_enforced_and_reported_on_another() that local share, while the shared counter was already over the cap. """ shared_redis = _SharedFakeRedis() - replica_a = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) - replica_b = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) + replica_a = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) + replica_b = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) key_hash = "vk-shared" model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "30d"}} user_api_key = UserAPIKeyAuth(token=key_hash, model_max_budget=model_max_budget) @@ -1436,7 +1436,7 @@ async def test_spend_logged_on_one_replica_is_enforced_and_reported_on_another() assert usage_on_b["gpt-4"]["current_spend"] == 1.25 # Control: a replica that never served this key reads the same total. - replica_c = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) + replica_c = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis)) with pytest.raises(litellm.BudgetExceededError): await replica_c.is_key_within_model_budget(user_api_key, "gpt-4") @@ -1459,7 +1459,7 @@ async def test_team_model_budget_is_shared_by_every_key_without_an_override(requ charge one team counter and are both refused once it is spent. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} check = lambda: limiter.is_team_within_model_budget( team_id="team-1", @@ -1507,7 +1507,7 @@ async def test_key_override_replaces_the_team_cap_for_that_model(): the team counter. """ dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} key_model_max_budget = {"gpt-4": {"budget_limit": 5.0, "time_period": "1d"}} await dual_cache.async_set_cache(key="team_model_spend:team-1:gpt-4:1d", value=9.0) @@ -1540,7 +1540,7 @@ async def test_key_override_replaces_the_team_cap_for_that_model(): async def test_key_entry_for_another_model_does_not_lift_the_team_cap(): """A key override only covers the model it names; other models stay on the team counter.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} key_model_max_budget = {"claude-3": {"budget_limit": 5.0, "time_period": "1d"}} @@ -1567,7 +1567,7 @@ async def test_key_entry_for_another_model_does_not_lift_the_team_cap(): @pytest.mark.asyncio async def test_team_budget_leaves_unconfigured_models_alone(): dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 0.0, "time_period": "1d"}} assert ( @@ -1593,7 +1593,7 @@ async def test_team_budget_leaves_unconfigured_models_alone(): async def test_team_counters_are_isolated_by_team_model_and_window(): """Same model on two teams, and two models with different windows on one team, never share a counter.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = { "gpt-4": {"budget_limit": 10.0, "time_period": "1d"}, "claude-3": {"budget_limit": 10.0, "time_period": "30d"}, @@ -1616,7 +1616,7 @@ async def test_team_counters_are_isolated_by_team_model_and_window(): @pytest.mark.asyncio async def test_malformed_team_entry_is_skipped_and_its_sibling_still_enforced(): - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) team_model_max_budget = { "gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}, "claude-3": {"budget_limit": 0.0, "time_period": "1d"}, @@ -1644,7 +1644,7 @@ async def test_malformed_team_entry_is_skipped_and_its_sibling_still_enforced(): async def test_malformed_key_entry_does_not_count_as_an_override(): """A key entry the limiter cannot enforce must not also switch the team cap off.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} key_model_max_budget = {"gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}} @@ -1680,7 +1680,7 @@ async def test_malformed_key_entry_does_not_count_as_an_override(): async def test_key_entry_without_a_spend_cap_does_not_lift_the_team_cap(key_entry): """A key row that only rate-limits the model, or has no enforceable cap, leaves the team cap in force.""" dual_cache = DualCache() - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache) team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}} key_model_max_budget = {"gpt-4": key_entry} diff --git a/tests/unit/proxy/management_endpoints/scim/test_scim_key_deactivation.py b/tests/unit/proxy/management_endpoints/scim/test_scim_key_deactivation.py index d35a676a28b..0c3600cbcc7 100644 --- a/tests/unit/proxy/management_endpoints/scim/test_scim_key_deactivation.py +++ b/tests/unit/proxy/management_endpoints/scim/test_scim_key_deactivation.py @@ -73,7 +73,7 @@ async def test_set_user_keys_blocked_flips_state_and_invalidates_cache(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(side_effect=fake_delete), ), ): @@ -104,7 +104,7 @@ async def test_set_user_keys_blocked_noop_when_no_matching_keys(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ) as mocked_delete, ): @@ -139,7 +139,7 @@ async def test_set_user_keys_unblocked_skips_admin_blocked_keys(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(side_effect=fake_delete), ), ): @@ -172,7 +172,7 @@ async def test_scim_delete_user_blocks_keys_before_deleting_user(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -220,7 +220,7 @@ async def test_scim_delete_user_clears_fk_referenced_rows_before_user_delete(): patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -294,7 +294,7 @@ async def test_scim_patch_user_active_false_blocks_keys(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -354,7 +354,7 @@ async def test_scim_patch_user_active_true_unblocks_keys(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -411,7 +411,7 @@ async def test_scim_patch_user_no_active_change_does_not_touch_keys(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -474,7 +474,7 @@ async def test_scim_put_user_omitting_active_preserves_deactivated_state(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): @@ -530,7 +530,7 @@ async def test_scim_put_user_explicit_active_false_blocks_keys(): AsyncMock(return_value=mock_scim_user), ), patch( - "litellm.proxy.management_endpoints.scim.scim_v2._delete_cache_key_object", + "litellm.proxy.management_endpoints.scim.scim_v2.delete_cache_key_object", AsyncMock(), ), ): diff --git a/tests/unit/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/unit/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index 62d77a00f25..88b28fffbd3 100644 --- a/tests/unit/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/unit/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -33,7 +33,7 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import ( _handle_group_membership_changes, _handle_team_membership_changes, _parse_member_entries, - _premium_user_check, + premium_user_check, _process_group_patch_operations, _recompute_scim_member_roles, _resolve_group_member_ids, @@ -493,7 +493,7 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp def scim_test_client(): """An in-process SCIM application with authorization dependencies bypassed.""" app = FastAPI() - app.dependency_overrides[_premium_user_check] = lambda: None + app.dependency_overrides[premium_user_check] = lambda: None app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) app.include_router(scim_router) return AsyncClient(transport=ASGITransport(app=app), base_url="http://test") diff --git a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py index 360fc5cc292..a6555faad98 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py @@ -3465,7 +3465,7 @@ async def test_member_billable_preview_checks_and_charges_destination_team( raise litellm.BudgetExceededError(current_cost=2, max_budget=1) checks: Final = AsyncMock(side_effect=check_and_tag) - monkeypatch.setattr(auth_module, "_run_centralized_common_checks", checks) + monkeypatch.setattr(auth_module, "run_centralized_common_checks", checks) http_request: Final = Request( { "type": "http", diff --git a/tests/unit/proxy/management_endpoints/test_cache_settings_endpoints.py b/tests/unit/proxy/management_endpoints/test_cache_settings_endpoints.py index 9a2dd914866..3721ab79bf1 100644 --- a/tests/unit/proxy/management_endpoints/test_cache_settings_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_cache_settings_endpoints.py @@ -201,11 +201,11 @@ async def test_update_cache_settings_persists_url_precedence(monkeypatch): mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = MagicMock() - proxy_config._encrypt_env_variables = MagicMock( + proxy_config.encrypt_env_variables = MagicMock( side_effect=lambda environment_variables: dict(environment_variables) ) - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) - proxy_config._init_cache = MagicMock() + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.init_cache = MagicMock() proxy_config.switch_on_llm_response_caching = MagicMock() with ( @@ -229,7 +229,7 @@ async def test_update_cache_settings_persists_url_precedence(monkeypatch): litellm_changed_by=None, ) - persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + persisted = proxy_config.encrypt_env_variables.call_args.kwargs["environment_variables"] assert persisted["url"] == "redis://:pw@host:6379/1" assert persisted["namespace"] == "ns" assert "host" not in persisted @@ -237,7 +237,7 @@ async def test_update_cache_settings_persists_url_precedence(monkeypatch): assert "db" not in persisted assert "password" not in persisted - init_params = proxy_config._init_cache.call_args.kwargs["cache_params"] + init_params = proxy_config.init_cache.call_args.kwargs["cache_params"] assert "host" not in init_params assert init_params["url"] == "redis://:pw@host:6379/1" @@ -262,7 +262,7 @@ async def test_get_cache_settings_masks_password_bearing_url(): mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -505,11 +505,11 @@ async def test_update_cache_settings_emits_audit_log_when_enabled(monkeypatch): mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = MagicMock() - proxy_config._encrypt_env_variables = MagicMock( + proxy_config.encrypt_env_variables = MagicMock( side_effect=lambda environment_variables: dict(environment_variables) ) - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) - proxy_config._init_cache = MagicMock() + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.init_cache = MagicMock() proxy_config.switch_on_llm_response_caching = MagicMock() audit_calls = [] @@ -575,11 +575,11 @@ async def test_update_cache_settings_no_audit_when_disabled(monkeypatch): mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = MagicMock() - proxy_config._encrypt_env_variables = MagicMock( + proxy_config.encrypt_env_variables = MagicMock( side_effect=lambda environment_variables: dict(environment_variables) ) - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) - proxy_config._init_cache = MagicMock() + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.init_cache = MagicMock() proxy_config.switch_on_llm_response_caching = MagicMock() audit_calls = [] @@ -802,7 +802,7 @@ async def test_get_cache_settings_falls_back_to_redis_env(monkeypatch): mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -829,7 +829,7 @@ async def test_get_cache_settings_redacts_password_with_marker(monkeypatch): mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -858,7 +858,7 @@ async def test_get_cache_settings_url_mode_hides_env_discrete_fields(monkeypatch mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -875,11 +875,11 @@ async def test_get_cache_settings_url_mode_hides_env_discrete_fields(monkeypatch def _mock_proxy_config_identity_crypto(): proxy_config = MagicMock() - proxy_config._encrypt_env_variables = MagicMock( + proxy_config.encrypt_env_variables = MagicMock( side_effect=lambda environment_variables: dict(environment_variables) ) - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) - proxy_config._init_cache = MagicMock() + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.init_cache = MagicMock() proxy_config.switch_on_llm_response_caching = MagicMock() return proxy_config @@ -913,7 +913,7 @@ async def test_update_preserves_stored_password_on_redacted_resubmit(monkeypatch litellm_changed_by=None, ) - persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + persisted = proxy_config.encrypt_env_variables.call_args.kwargs["environment_variables"] assert persisted["host"] == "oldhost" assert persisted["namespace"] == "edited" assert persisted["password"] == "realpw" @@ -945,7 +945,7 @@ async def test_update_drops_env_sourced_redacted_secret(monkeypatch): litellm_changed_by=None, ) - persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + persisted = proxy_config.encrypt_env_variables.call_args.kwargs["environment_variables"] assert "password" not in persisted @@ -974,7 +974,7 @@ async def test_update_applies_new_password(monkeypatch): litellm_changed_by=None, ) - persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + persisted = proxy_config.encrypt_env_variables.call_args.kwargs["environment_variables"] assert persisted["password"] == "brandnewpw" @@ -992,7 +992,7 @@ async def test_test_cache_connection_survives_saved_lookup_failure(monkeypatch): # a client whose find_unique is not awaitable, so the saved read raises bad_prisma = MagicMock() proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) cache_instance = MagicMock() cache_instance.cache = MagicMock() @@ -1029,7 +1029,7 @@ async def test_get_cache_settings_does_not_surface_non_display_env_credentials(m mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -1058,7 +1058,7 @@ async def test_test_cache_connection_does_not_log_plaintext_credentials(monkeypa mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) cache_instance = MagicMock() cache_instance.cache = MagicMock() @@ -1096,7 +1096,7 @@ async def test_test_cache_connection_does_not_replay_saved_password_to_new_host( mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) proxy_config = MagicMock() - proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config.decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) cache_instance = MagicMock() cache_instance.cache = MagicMock() diff --git a/tests/unit/proxy/management_endpoints/test_common_utils.py b/tests/unit/proxy/management_endpoints/test_common_utils.py index 67920c9c7fe..43c74286d75 100644 --- a/tests/unit/proxy/management_endpoints/test_common_utils.py +++ b/tests/unit/proxy/management_endpoints/test_common_utils.py @@ -28,12 +28,12 @@ from litellm.proxy._types import ( from litellm.proxy.management_endpoints.common_utils import ( _has_non_empty_value, _org_admin_can_invite_user, - _set_object_metadata_field, _team_admin_can_invite_user, - _update_metadata_fields, - _user_has_admin_privileges, - _user_has_admin_view, admin_can_invite_user, + set_object_metadata_field, + update_metadata_fields, + user_api_key_has_admin_view, + user_has_admin_privileges, ) from litellm.types.utils import BudgetConfig @@ -53,7 +53,7 @@ class TestUpdateMetadataFieldsEmptyCollections: guardrails by sending `guardrails: []`). """ - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_empty_list_does_not_trigger_premium_check(self, mock_premium_check): """Empty lists for premium fields must not trigger the premium check.""" updated_kv = { @@ -62,10 +62,10 @@ class TestUpdateMetadataFieldsEmptyCollections: "policies": [], "logging": [], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_empty_list_still_updates_metadata(self, mock_premium_check): """ Empty lists must still be moved into metadata so users can clear @@ -76,7 +76,7 @@ class TestUpdateMetadataFieldsEmptyCollections: "guardrails": [], "policies": [], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) # The fields should have been moved into metadata assert ( "guardrails" not in updated_kv @@ -85,17 +85,17 @@ class TestUpdateMetadataFieldsEmptyCollections: assert updated_kv["metadata"]["guardrails"] == [] assert updated_kv["metadata"]["policies"] == [] - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_empty_dict_does_not_trigger_premium_check(self, mock_premium_check): """Empty dicts for premium fields must not trigger the premium check.""" updated_kv = { "team_id": "test-team", "secret_manager_settings": {}, } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_empty_dict_still_updates_metadata(self, mock_premium_check): """ Empty dicts must still be moved into metadata so users can clear @@ -105,13 +105,13 @@ class TestUpdateMetadataFieldsEmptyCollections: "team_id": "test-team", "secret_manager_settings": {}, } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) assert ( "secret_manager_settings" not in updated_kv ), "secret_manager_settings should be popped from top-level" assert updated_kv["metadata"]["secret_manager_settings"] == {} - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_none_value_does_not_trigger_premium_check(self, mock_premium_check): """None values for premium fields should be silently ignored.""" updated_kv = { @@ -119,51 +119,51 @@ class TestUpdateMetadataFieldsEmptyCollections: "guardrails": None, "policies": None, } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_absent_fields_do_not_trigger_premium_check(self, mock_premium_check): """Fields not present in the dict should not trigger premium check.""" updated_kv = { "team_id": "test-team", "team_alias": "example-team", } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_non_empty_list_triggers_premium_check(self, mock_premium_check): """Non-empty lists for premium fields should trigger the premium check.""" updated_kv = { "team_id": "test-team", "guardrails": ["my-guardrail"], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_non_empty_value_triggers_premium_check(self, mock_premium_check): """Non-empty string values for premium fields should trigger the premium check.""" updated_kv = { "team_id": "test-team", "tags": ["production"], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_non_empty_list_updates_metadata(self, mock_premium_check): """Non-empty lists should be moved into metadata.""" updated_kv = { "team_id": "test-team", "guardrails": ["my-guardrail"], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) assert "guardrails" not in updated_kv assert updated_kv["metadata"]["guardrails"] == ["my-guardrail"] - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_false_boolean_does_not_trigger_premium_check(self, mock_premium_check): """ Regression #30285: /team/update sends disable_global_guardrails=False @@ -171,25 +171,25 @@ class TestUpdateMetadataFieldsEmptyCollections: premium check, so non-premium users are not wrongly 403'd. """ updated_kv = {"team_id": "test-team", "disable_global_guardrails": False} - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_false_boolean_still_updates_metadata(self, mock_premium_check): """A falsy boolean must still be moved into metadata so it persists.""" updated_kv = {"team_id": "test-team", "disable_global_guardrails": False} - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) assert "disable_global_guardrails" not in updated_kv assert updated_kv["metadata"]["disable_global_guardrails"] is False - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_true_boolean_triggers_premium_check(self, mock_premium_check): """Control: enabling the premium feature (True) still requires a license.""" updated_kv = {"team_id": "test-team", "disable_global_guardrails": True} - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_called() - @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + @patch("litellm.proxy.management_endpoints.common_utils.premium_user_check") def test_ui_typical_payload_does_not_trigger_premium_check( self, mock_premium_check ): @@ -208,7 +208,7 @@ class TestUpdateMetadataFieldsEmptyCollections: }, "policies": [], } - _update_metadata_fields(updated_kv=updated_kv) + update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() @@ -228,7 +228,7 @@ class TestUserHasAdminView: """Parametrized test: admin roles return True, non-admin return False.""" mock_auth = MagicMock() mock_auth.user_role = user_role - assert _user_has_admin_view(mock_auth) == expected + assert user_api_key_has_admin_view(mock_auth) == expected def test_user_has_admin_view_with_user_api_key_auth(self): """Test with actual UserAPIKeyAuth object.""" @@ -242,8 +242,8 @@ class TestUserHasAdminView: api_key="sk-yyy", user_role=LitellmUserRoles.INTERNAL_USER, ) - assert _user_has_admin_view(auth_admin) is True - assert _user_has_admin_view(auth_user) is False + assert user_api_key_has_admin_view(auth_admin) is True + assert user_api_key_has_admin_view(auth_user) is False def test_published_enterprise_import_of_team_admin_check_still_answers(): @@ -384,7 +384,7 @@ class TestUserHasAdminPrivileges: api_key="sk-x", user_role=LitellmUserRoles.PROXY_ADMIN, ) - result = await _user_has_admin_privileges( + result = await user_has_admin_privileges( user_api_key_dict=auth, prisma_client=None, ) @@ -398,7 +398,7 @@ class TestUserHasAdminPrivileges: api_key="sk-x", user_role=LitellmUserRoles.INTERNAL_USER, ) - result = await _user_has_admin_privileges( + result = await user_has_admin_privileges( user_api_key_dict=auth, prisma_client=None, ) @@ -455,9 +455,9 @@ class TestSetObjectMetadataField: """Parametrized test: premium fields trigger _premium_user_check.""" team = LiteLLM_TeamTable(team_id="t1", metadata={}) with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ) as mock_premium: - _set_object_metadata_field(team, field_name, value) + set_object_metadata_field(team, field_name, value) if should_call_premium: mock_premium.assert_called_once() else: @@ -468,9 +468,9 @@ class TestSetObjectMetadataField: """Test initializes metadata dict when object has None.""" team = LiteLLM_TeamTable(team_id="t1", metadata=None) with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ): - _set_object_metadata_field(team, "model_rpm_limit", {"x": 1}) + set_object_metadata_field(team, "model_rpm_limit", {"x": 1}) assert team.metadata == {"model_rpm_limit": {"x": 1}} def test_mcp_rpm_limit_is_hoisted_into_metadata(self): @@ -492,11 +492,11 @@ class TestSetObjectMetadataField: data = SimpleNamespace(mcp_rpm_limit=mcp_rpm_limit) with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ): for field in LiteLLM_ManagementEndpoint_MetadataFields: if getattr(data, field, None) is not None: - _set_object_metadata_field(team, field, getattr(data, field)) + set_object_metadata_field(team, field, getattr(data, field)) assert team.metadata["mcp_rpm_limit"] == mcp_rpm_limit @@ -681,7 +681,7 @@ class TestCheckPassthroughRoutesCallerPermission: from pydantic import BaseModel from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) class _RouteData(BaseModel): @@ -690,7 +690,7 @@ class TestCheckPassthroughRoutesCallerPermission: data = _RouteData(allowed_passthrough_routes=["/v1/foo"]) with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission(data, self._non_admin()) + check_passthrough_routes_caller_permission(data, self._non_admin()) assert exc_info.value.status_code == 403 assert exc_info.value.detail == { @@ -702,7 +702,7 @@ class TestCheckPassthroughRoutesCallerPermission: from pydantic import BaseModel from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) class _RouteData(BaseModel): @@ -711,7 +711,7 @@ class TestCheckPassthroughRoutesCallerPermission: data = _RouteData(metadata={"allowed_passthrough_routes": ["/v1/foo"]}) with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission(data, self._non_admin()) + check_passthrough_routes_caller_permission(data, self._non_admin()) assert exc_info.value.detail == { "error": "Only proxy admins can set `metadata.allowed_passthrough_routes` on a key." @@ -721,13 +721,13 @@ class TestCheckPassthroughRoutesCallerPermission: from pydantic import BaseModel from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) class _Bare(BaseModel): unrelated: str = "x" - assert _check_passthrough_routes_caller_permission(_Bare(), self._non_admin()) is None + assert check_passthrough_routes_caller_permission(_Bare(), self._non_admin()) is None @pytest.mark.parametrize( "kwargs, field", @@ -741,7 +741,7 @@ class TestCheckPassthroughRoutesCallerPermission: from pydantic import BaseModel from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) class _RouteData(BaseModel): @@ -749,7 +749,7 @@ class TestCheckPassthroughRoutesCallerPermission: metadata: dict[str, object] | None = None with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _RouteData.model_validate(kwargs), self._non_admin(), entity="team" ) @@ -777,10 +777,10 @@ class TestDeniedPassthroughRoutesCallerPermission: ids=["cleared", "replaced", "dropped-by-metadata-replace", "dropped-by-null-metadata"], ) def test_non_admin_cannot_change_an_existing_deny_list(self, kwargs: dict[str, object], field: str) -> None: - from litellm.proxy.management_endpoints.common_utils import _check_passthrough_routes_caller_permission + from litellm.proxy.management_endpoints.common_utils import check_passthrough_routes_caller_permission with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _DenyRouteData.model_validate(kwargs), UserAPIKeyAuth(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER), existing_metadata=_EXISTING_DENY, @@ -799,28 +799,28 @@ class TestDeniedPassthroughRoutesCallerPermission: ids=["resent-top-level", "resent-in-metadata", "unrelated-field"], ) def test_non_admin_may_leave_an_existing_deny_list_unchanged(self, kwargs: dict[str, object]) -> None: - from litellm.proxy.management_endpoints.common_utils import _check_passthrough_routes_caller_permission + from litellm.proxy.management_endpoints.common_utils import check_passthrough_routes_caller_permission - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _DenyRouteData.model_validate(kwargs), UserAPIKeyAuth(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER), existing_metadata=_EXISTING_DENY, ) def test_non_admin_may_send_null_metadata_when_no_deny_list_exists(self) -> None: - from litellm.proxy.management_endpoints.common_utils import _check_passthrough_routes_caller_permission + from litellm.proxy.management_endpoints.common_utils import check_passthrough_routes_caller_permission - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _DenyRouteData(metadata=None), UserAPIKeyAuth(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER), existing_metadata={"team": "core"}, ) def test_malformed_metadata_deny_entries_are_rejected_even_for_proxy_admins(self) -> None: - from litellm.proxy.management_endpoints.common_utils import _check_passthrough_routes_caller_permission + from litellm.proxy.management_endpoints.common_utils import check_passthrough_routes_caller_permission with pytest.raises(HTTPException) as exc_info: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( _DenyRouteData(metadata={"denied_passthrough_routes": [123, None]}), UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), ) @@ -849,11 +849,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission(True, None, self._non_admin()) + check_disable_global_guardrails_caller_permission(True, None, self._non_admin()) assert exc_info.value.status_code == 403 assert exc_info.value.detail == {"error": "Only proxy admins can set `disable_global_guardrails` on a key."} @@ -862,11 +862,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( None, {"disable_global_guardrails": True}, self._non_admin() ) @@ -877,11 +877,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( False, {"disable_global_guardrails": True}, self._non_admin() ) @@ -892,37 +892,37 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission(True, None, self._non_admin(), entity="team") + check_disable_global_guardrails_caller_permission(True, None, self._non_admin(), entity="team") assert exc_info.value.detail == {"error": "Only proxy admins can set `disable_global_guardrails` on a team."} def test_false_and_absent_flag_do_not_raise(self): from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) non_admin = self._non_admin() - assert _check_disable_global_guardrails_caller_permission(False, None, non_admin) is None - assert _check_disable_global_guardrails_caller_permission(None, None, non_admin) is None - assert _check_disable_global_guardrails_caller_permission(None, {}, non_admin) is None + assert check_disable_global_guardrails_caller_permission(False, None, non_admin) is None + assert check_disable_global_guardrails_caller_permission(None, None, non_admin) is None + assert check_disable_global_guardrails_caller_permission(None, {}, non_admin) is None assert ( - _check_disable_global_guardrails_caller_permission(None, {"disable_global_guardrails": False}, non_admin) + check_disable_global_guardrails_caller_permission(None, {"disable_global_guardrails": False}, non_admin) is None ) def test_unchanged_stored_flag_does_not_raise(self): """Re-sending a flag that is already stored is not an opt-out.""" from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) non_admin = self._non_admin() assert ( - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( True, {"disable_global_guardrails": True}, non_admin, @@ -935,11 +935,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: from fastapi import HTTPException from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) with pytest.raises(HTTPException) as exc_info: - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( True, None, self._non_admin(), @@ -951,11 +951,11 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: def test_proxy_admin_may_set_the_flag(self): from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) assert ( - _check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, self._admin()) + check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, self._admin()) is None ) @@ -963,7 +963,7 @@ class TestCheckDisableGlobalGuardrailsCallerPermission: class TestTeamMemberHasPermission: def test_requires_caller_to_be_a_team_member(self): from litellm.proxy.management_endpoints.common_utils import ( - _team_member_has_permission, + team_member_has_permission, ) team = LiteLLM_TeamTable( @@ -974,7 +974,7 @@ class TestTeamMemberHasPermission: key = UserAPIKeyAuth( user_id="u1", api_key="sk-x", user_role=LitellmUserRoles.INTERNAL_USER ) - assert _team_member_has_permission(key, team, "/key/generate") is False + assert team_member_has_permission(key, team, "/key/generate") is False class TestUserHasAdminPrivilegesGuard: @@ -986,7 +986,7 @@ class TestUserHasAdminPrivilegesGuard: ) mock_get_user = AsyncMock(return_value=None) with patch("litellm.proxy.auth.auth_checks.get_user_object", mock_get_user): - result = await _user_has_admin_privileges( + result = await user_has_admin_privileges( user_api_key_dict=auth, prisma_client=None ) assert result is False @@ -1013,7 +1013,7 @@ class TestUserHasAdminPrivilegesGuard: ) mock_get_user = AsyncMock(return_value=user_obj) with patch("litellm.proxy.auth.auth_checks.get_user_object", mock_get_user): - result = await _user_has_admin_privileges( + result = await user_has_admin_privileges( user_api_key_dict=auth, prisma_client=MagicMock() ) assert result is True @@ -1104,9 +1104,9 @@ class TestSetObjectMetadataFieldPremiumArg: def test_premium_check_receives_the_field_name(self): team = LiteLLM_TeamTable(team_id="t1", metadata={}) with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ) as mock_premium: - _set_object_metadata_field(team, "guardrails", ["g1"]) + set_object_metadata_field(team, "guardrails", ["g1"]) mock_premium.assert_called_once_with("guardrails") @@ -1124,9 +1124,9 @@ class TestUpdateMetadataFieldMove: def test_set_premium_field_is_moved_into_metadata(self): updated_kv = {"guardrails": ["g1"]} with patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check" + "litellm.proxy.management_endpoints.common_utils.premium_user_check" ): - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) assert "guardrails" not in updated_kv assert updated_kv["metadata"]["guardrails"] == ["g1"] @@ -1171,7 +1171,7 @@ class TestUpdateMetadataFieldsPremiumCheck: """ @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", side_effect=Exception("Should not be called"), ) def test_empty_policies_skips_premium_check(self, mock_check): @@ -1181,11 +1181,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_alias": "my-team", "policies": [], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_not_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", side_effect=Exception("Should not be called"), ) def test_empty_guardrails_skips_premium_check(self, mock_check): @@ -1194,11 +1194,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "guardrails": [], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_not_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", side_effect=Exception("Should not be called"), ) def test_empty_string_team_member_key_duration_skips_premium_check( @@ -1209,11 +1209,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "team_member_key_duration": "", } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_not_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", side_effect=Exception("Should not be called"), ) def test_full_ui_payload_with_empty_premium_fields_skips_premium_check( @@ -1231,11 +1231,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_member_key_duration": "", "prompts": [], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_not_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", ) def test_non_empty_policies_triggers_premium_check(self, mock_check): """policies: ['real-policy'] SHOULD trigger premium user check.""" @@ -1243,11 +1243,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "policies": ["real-policy"], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", ) def test_non_empty_guardrails_triggers_premium_check(self, mock_check): """guardrails: ['my-guardrail'] SHOULD trigger premium user check.""" @@ -1255,11 +1255,11 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "guardrails": ["my-guardrail"], } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_called() @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", + "litellm.proxy.management_endpoints.common_utils.premium_user_check", ) def test_non_empty_team_member_key_duration_triggers_premium_check( self, mock_check @@ -1269,7 +1269,7 @@ class TestUpdateMetadataFieldsPremiumCheck: "team_id": "team-123", "team_member_key_duration": "30d", } - _update_metadata_fields(updated_kv) + update_metadata_fields(updated_kv) mock_check.assert_called() diff --git a/tests/unit/proxy/management_endpoints/test_config_override_endpoints.py b/tests/unit/proxy/management_endpoints/test_config_override_endpoints.py index 49b0ed1b28a..8d736bfdbdd 100644 --- a/tests/unit/proxy/management_endpoints/test_config_override_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_config_override_endpoints.py @@ -14,7 +14,7 @@ from litellm.proxy.management_endpoints.config_override_endpoints import ( CYBERARK_ENV_VAR_MAPPING, HASHICORP_ENV_VAR_MAPPING, _build_field_schema, - _set_env_vars, + set_env_vars, ) from litellm.proxy.proxy_server import app from litellm.types.proxy.management_endpoints.config_overrides import ( @@ -54,6 +54,14 @@ def _make_mock_proxy_config(): k: v.replace("enc_", "") if isinstance(v, str) else v for k, v in d.items() } ) + cfg.encrypt_env_variables = MagicMock( + side_effect=lambda d: {k: f"enc_{v}" for k, v in d.items()} + ) + cfg.decrypt_db_variables = MagicMock( + side_effect=lambda d: { + k: v.replace("enc_", "") if isinstance(v, str) else v for k, v in d.items() + } + ) return cfg @@ -194,7 +202,7 @@ async def test_hashicorp_vault_crud_lifecycle(client, monkeypatch): # 10. _set_env_vars: empty string unsets monkeypatch.setenv("HCP_VAULT_TOKEN", "existing") - _set_env_vars({"vault_token": "", "vault_addr": "https://v.com"}) + set_env_vars({"vault_token": "", "vault_addr": "https://v.com"}) assert os.environ.get("HCP_VAULT_TOKEN") is None assert os.environ["HCP_VAULT_ADDR"] == "https://v.com" @@ -209,9 +217,9 @@ async def test_hashicorp_vault_crud_lifecycle(client, monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key") pc = ProxyConfig() orig = {"vault_addr": "https://v.com", "vault_token": "secret"} - encrypted = pc._encrypt_env_variables(orig) + encrypted = pc.encrypt_env_variables(orig) assert all(encrypted[k] != orig[k] for k in orig) - decrypted = pc._decrypt_db_variables(encrypted) + decrypted = pc.decrypt_db_variables(encrypted) assert all(decrypted[k] == orig[k] for k in orig) finally: diff --git a/tests/unit/proxy/management_endpoints/test_coordination_redis_endpoints.py b/tests/unit/proxy/management_endpoints/test_coordination_redis_endpoints.py index 288d8847fbb..2ee7dd1364a 100644 --- a/tests/unit/proxy/management_endpoints/test_coordination_redis_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_coordination_redis_endpoints.py @@ -191,7 +191,7 @@ async def test_get_source_does_not_build_a_client(monkeypatch): with ( patch("litellm.proxy.proxy_server.prisma_client", _prisma_with_general_settings({})), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache") as mock_build, + patch("litellm.proxy.proxy_server.build_redis_usage_cache") as mock_build, ): response = await get_coordination_redis_settings(user_api_key_dict=_admin_auth()) @@ -480,7 +480,7 @@ async def test_connection_test_returns_healthy_on_successful_ping(): with ( patch("litellm.proxy.proxy_server.prisma_client", _prisma_with_general_settings({})), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache", return_value=mock_client) as mock_build, + patch("litellm.proxy.proxy_server.build_redis_usage_cache", return_value=mock_client) as mock_build, ): response = await check_coordination_redis_connection( request=CoordinationRedisSettingsRequest( @@ -509,7 +509,7 @@ async def test_connection_test_reports_unhealthy_without_leaking_the_password(): with ( patch("litellm.proxy.proxy_server.prisma_client", _prisma_with_general_settings({})), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache", return_value=mock_client), + patch("litellm.proxy.proxy_server.build_redis_usage_cache", return_value=mock_client), ): response = await check_coordination_redis_connection( request=CoordinationRedisSettingsRequest( @@ -543,7 +543,7 @@ async def test_connection_test_uses_the_saved_password_for_a_redacted_field(): _prisma_with_general_settings({"coordination_redis": _SAVED_SETTINGS}), ), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache", return_value=mock_client) as mock_build, + patch("litellm.proxy.proxy_server.build_redis_usage_cache", return_value=mock_client) as mock_build, ): response = await check_coordination_redis_connection( request=CoordinationRedisSettingsRequest( @@ -568,7 +568,7 @@ async def test_connection_test_times_out_instead_of_hanging(): with ( patch("litellm.proxy.proxy_server.prisma_client", _prisma_with_general_settings({})), patch("litellm.proxy.proxy_server.proxy_config", _proxy_config()), - patch("litellm.proxy.proxy_server._build_redis_usage_cache", return_value=mock_client), + patch("litellm.proxy.proxy_server.build_redis_usage_cache", return_value=mock_client), patch( "litellm.proxy.management_endpoints.coordination_redis_endpoints._PING_TIMEOUT_SECONDS", 0.01, diff --git a/tests/unit/proxy/management_endpoints/test_credential_migration.py b/tests/unit/proxy/management_endpoints/test_credential_migration.py index ac5a45499d3..ffb855173fa 100644 --- a/tests/unit/proxy/management_endpoints/test_credential_migration.py +++ b/tests/unit/proxy/management_endpoints/test_credential_migration.py @@ -17,7 +17,7 @@ import pytest from litellm._service_logger import ServiceTypes from litellm.proxy import proxy_server from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - _V2_GCM_PREFIX, + V2_GCM_PREFIX, encrypt_value_helper, ) from litellm.proxy.management_endpoints import credential_migration as cm @@ -82,7 +82,7 @@ def test_reencrypt_value_legacy_to_v2(salt_key, monkeypatch): out = cm.reencrypt_value(legacy) assert out != legacy - assert out.startswith(_V2_GCM_PREFIX) + assert out.startswith(V2_GCM_PREFIX) def test_reencrypt_value_is_idempotent(salt_key, monkeypatch): @@ -113,7 +113,7 @@ def test_reencrypt_selective_dict(salt_key, monkeypatch): data = {"api_key": legacy_key, "base_url": "https://x", "integration_token": None} out = cm.reencrypt_selective_dict(data, ["api_key", "integration_token"]) - assert out["api_key"].startswith(_V2_GCM_PREFIX) + assert out["api_key"].startswith(V2_GCM_PREFIX) assert out["base_url"] == "https://x" # untouched non-sensitive assert out["integration_token"] is None # null skipped @@ -164,7 +164,7 @@ async def test_vantage_walker_migrates_legacy_field(salt_key, monkeypatch): written = json.loads( client.db.litellm_config.update.call_args.kwargs["data"]["param_value"] ) - assert written["api_key"].startswith(_V2_GCM_PREFIX) + assert written["api_key"].startswith(V2_GCM_PREFIX) assert written["base_url"] == "https://api.vantage.sh" # non-sensitive untouched @@ -584,11 +584,11 @@ async def test_migrate_covered_tables_reports_real_counts(salt_key, monkeypatch) client.db.litellm_config.find_unique = AsyncMock(return_value=None) async def fake_rotate(**kwargs): - # Stand in for _rotate_master_key: re-encrypt the model api_key in place. + # Stand in for rotate_master_key: re-encrypt the model api_key in place. row.litellm_params["api_key"] = v2 monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._rotate_master_key", + "litellm.proxy.management_endpoints.key_management_endpoints.rotate_master_key", fake_rotate, ) diff --git a/tests/unit/proxy/management_endpoints/test_delete_verification_tokens_failed.py b/tests/unit/proxy/management_endpoints/test_delete_verification_tokens_failed.py index e33945df7dc..bfd0d663eac 100644 --- a/tests/unit/proxy/management_endpoints/test_delete_verification_tokens_failed.py +++ b/tests/unit/proxy/management_endpoints/test_delete_verification_tokens_failed.py @@ -94,7 +94,7 @@ async def test_delete_all_tokens_admin_returns_empty_failed_tokens(monkeypatch): mock_cache.delete_cache = MagicMock() monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( @@ -132,7 +132,7 @@ async def test_delete_tokens_non_admin_all_succeed_returns_empty_failed_tokens( mock_cache.delete_cache = MagicMock() monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( @@ -183,7 +183,7 @@ async def test_delete_tokens_non_admin_token_not_in_db_returns_failed_tokens( mock_cache.delete_cache = MagicMock() monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( @@ -234,7 +234,7 @@ async def test_delete_tokens_admin_partial_db_failure_returns_failed_tokens( mock_cache.delete_cache = MagicMock() monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( diff --git a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index 712107e4a32..39265b37b95 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -4004,7 +4004,7 @@ async def test_admin_user_update_spend_invalidates_counter(mocker): mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") mock_invalidate = mocker.patch( - "litellm.proxy.proxy_server._invalidate_spend_counter", + "litellm.proxy.proxy_server.invalidate_spend_counter", new=mocker.AsyncMock(), ) @@ -4038,7 +4038,7 @@ async def test_user_update_rejects_non_finite_spend(mocker): mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") mock_invalidate = mocker.patch( - "litellm.proxy.proxy_server._invalidate_spend_counter", + "litellm.proxy.proxy_server.invalidate_spend_counter", new=mocker.AsyncMock(), ) @@ -4276,7 +4276,7 @@ def _object_permission_mocks(mocker, existing_object_permission_id=None): mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") mocker.patch( - "litellm.proxy.proxy_server._invalidate_spend_counter", + "litellm.proxy.proxy_server.invalidate_spend_counter", new=mocker.AsyncMock(), ) return mock_prisma_client diff --git a/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py b/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py index a4126c476bd..3d9667ace1b 100644 --- a/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py +++ b/tests/unit/proxy/management_endpoints/test_key_generate_prisma.py @@ -514,9 +514,9 @@ def test_call_with_user_over_budget(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -612,9 +612,9 @@ def test_call_with_end_user_over_budget(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -722,9 +722,9 @@ def test_call_with_proxy_over_budget(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -814,9 +814,9 @@ def test_call_with_user_over_budget_stream(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -921,9 +921,9 @@ def test_call_with_proxy_over_budget_stream(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() resp = ModelResponse( id="chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac", @@ -1573,9 +1573,9 @@ def test_call_with_key_over_budget(prisma_client): # update spend using track_cost callback, make 2nd request, it should fail from litellm import Choices, Message, ModelResponse, Usage from litellm.caching.caching import Cache - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() litellm.cache = Cache() import time @@ -1690,7 +1690,7 @@ def test_call_with_key_over_budget_no_cache(prisma_client): print("result from user auth with new key", result) # update spend using track_cost callback, make 2nd request, it should fail - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger from litellm.proxy.proxy_server import user_api_key_cache user_api_key_cache.in_memory_cache.cache_dict = {} @@ -1720,7 +1720,7 @@ def test_call_with_key_over_budget_no_cache(prisma_client): model="gpt-35-turbo", # azure always has model written like this usage=Usage(prompt_tokens=210, completion_tokens=200, total_tokens=410), ) - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() await proxy_db_logger._PROXY_track_cost_callback( kwargs={ "model": "chatgpt-v-3", @@ -1943,9 +1943,9 @@ async def test_call_with_key_never_over_budget(prisma_client): from litellm._uuid import uuid from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() request_id = f"chatcmpl-{uuid.uuid4()}" @@ -2034,9 +2034,9 @@ async def test_call_with_key_over_budget_stream(prisma_client): from litellm._uuid import uuid from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() request_id = f"chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac{uuid.uuid4()}" resp = ModelResponse( @@ -2346,31 +2346,31 @@ async def test_upperbound_key_param_none_duration(prisma_client): def test_get_bearer_token(): - from litellm.proxy.auth.user_api_key_auth import _get_bearer_token + from litellm.proxy.auth.user_api_key_auth import get_bearer_token # Test valid Bearer token api_key = "Bearer valid_token" - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "valid_token", f"Expected 'valid_token', got '{result}'" # Test empty API key api_key = "" - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "", f"Expected '', got '{result}'" # Test API key without Bearer prefix api_key = "invalid_token" - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "", f"Expected '', got '{result}'" # Test API key with Bearer prefix and extra spaces api_key = " Bearer valid_token " - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "", f"Expected '', got '{result}'" # Test API key with Bearer prefix and no token api_key = "Bearer sk-9876" - result = _get_bearer_token(api_key) + result = get_bearer_token(api_key) assert result == "sk-9876", f"Expected 'sk-9876', got '{result}'" @@ -2507,7 +2507,7 @@ async def track_cost_callback_helper_fn(generated_key: str, user_id: str): from litellm._uuid import uuid from litellm import Choices, Message, ModelResponse, Usage - from litellm.proxy.proxy_server import _ProxyDBLogger + from litellm.proxy.proxy_server import ProxyDBLogger request_id = f"chatcmpl-e41836bb-bb8b-4df2-8e70-8f3e160155ac{uuid.uuid4()}" resp = ModelResponse( @@ -2525,7 +2525,7 @@ async def track_cost_callback_helper_fn(generated_key: str, user_id: str): model="gpt-35-turbo", # azure always has model written like this usage=Usage(prompt_tokens=210, completion_tokens=200, total_tokens=410), ) - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() await proxy_db_logger._PROXY_track_cost_callback( kwargs={ "call_type": "acompletion", diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index 20672e67358..0f21f3ee3e0 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -38,7 +38,7 @@ from litellm.proxy._types import ( ) from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.proxy.auth.auth_checks import ( - _delete_cache_key_object, + delete_cache_key_object, jwt_key_mapping_cache_key, ) from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth @@ -55,8 +55,8 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( _enforce_upperbound_key_params, _execute_virtual_key_regeneration, _get_and_validate_existing_key, - _list_key_helper, - _persist_deleted_verification_tokens, + list_key_helper, + persist_deleted_verification_tokens, _process_single_key_update, _requested_end_user_budget_id, _save_deleted_verification_token_records, @@ -108,7 +108,7 @@ async def test_list_keys(): "admin_team_ids": ["28bd3181-02c5-48f2-b408-ce790fb3d5ba"], } try: - result = await _list_key_helper(**args) + result = await list_key_helper(**args) except Exception as e: print(f"error: {e}") @@ -155,7 +155,7 @@ async def test_list_keys_include_created_by_keys(): } try: - result = await _list_key_helper(**args) + result = await list_key_helper(**args) except Exception as e: print(f"error: {e}") @@ -219,7 +219,7 @@ async def test_list_keys_include_created_by_keys(): args["include_created_by_keys"] = False try: - result = await _list_key_helper(**args) + result = await list_key_helper(**args) except Exception as e: print(f"error: {e}") @@ -246,7 +246,7 @@ async def test_list_keys_include_created_by_keys(): ) try: - result = await _list_key_helper(**args) + result = await list_key_helper(**args) except Exception as e: print(f"error: {e}") @@ -1468,7 +1468,7 @@ async def test_list_keys_full_object_returns_lifetime_total_spend(): ) mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=1) - result = await _list_key_helper( + result = await list_key_helper( prisma_client=mock_prisma_client, page=1, size=50, @@ -2708,7 +2708,7 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch): pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) @@ -3066,7 +3066,7 @@ async def test_update_key_by_alias_only(monkeypatch): request_data = UpdateKeyRequest(key_alias="prod-alias", max_budget=50.0) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache: mock_delete_cache.return_value = None result = await update_key_fn( @@ -3121,7 +3121,7 @@ async def test_update_key_changed_alias_must_match_key_alias_pattern(monkeypatch mock_prisma_client.update_data.assert_not_awaited() with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", return_value=None, ): await update_key_fn( @@ -3237,7 +3237,7 @@ async def test_update_key_with_key_and_alias_selects_by_key(monkeypatch): ) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache: mock_delete_cache.return_value = None result = await update_key_fn( @@ -3307,7 +3307,7 @@ async def test_block_key_existing_key_succeeds(monkeypatch): pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) @@ -5379,7 +5379,7 @@ async def test_persist_deleted_verification_tokens(): allowed_routes=[], ) - await _persist_deleted_verification_tokens( + await persist_deleted_verification_tokens( keys=[key], prisma_client=mock_prisma_client, user_api_key_dict=user_api_key_dict, @@ -5467,7 +5467,7 @@ async def test_delete_verification_tokens_persists_deleted_keys(monkeypatch): return token if not token.startswith("sk-") else f"hashed-{token}" monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", mock_hash_token, ) monkeypatch.setattr( @@ -5580,7 +5580,7 @@ async def test_delete_verification_tokens_evicts_jwt_key_mapping_cache(monkeypat recording_evict, ) monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", lambda token: token, ) monkeypatch.setattr( @@ -6150,7 +6150,7 @@ async def test_list_keys_with_expand_user(): "expand": ["user"], # Test the expand parameter } - result = await _list_key_helper(**args) + result = await list_key_helper(**args) # Verify that keys were fetched mock_find_many_keys.assert_called_once() @@ -6261,7 +6261,7 @@ async def test_list_keys_with_expand_user_includes_created_by_user(): "expand": ["user"], } - result = await _list_key_helper(**args) + result = await list_key_helper(**args) # Verify that the user lookup included both user_id and created_by call_args = mock_find_many_users.call_args @@ -6342,7 +6342,7 @@ async def test_list_keys_with_status_deleted(): "status": "deleted", # Test the status parameter } - result = await _list_key_helper(**args) + result = await list_key_helper(**args) # Verify that deleted table was queried mock_find_many_deleted.assert_called_once() @@ -6481,7 +6481,7 @@ async def test_list_key_helper_revoked_status_filters_live_table_on_blocked(): mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) mock_prisma_client.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) - await _list_key_helper( + await list_key_helper( prisma_client=mock_prisma_client, page=1, size=50, @@ -6676,7 +6676,7 @@ async def test_list_keys_non_admin_user_id_auto_set(): return_value=[], ): with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", mock_list_key_helper, ): mock_request = Mock() @@ -6755,7 +6755,7 @@ async def _invoke_list_keys_and_capture_helper_kwargs( AsyncMock(return_value=team_objects), ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", mock_list_key_helper, ), ): @@ -7176,7 +7176,7 @@ async def test_list_key_helper_applies_search_to_prisma_where(): mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) - await _list_key_helper( + await list_key_helper( prisma_client=mock_prisma_client, page=1, size=50, @@ -7219,7 +7219,7 @@ async def _run_bulk_update_on_one_key( with ( patch( # test-quality-ok: the handler reads the cache and hook singletons from module globals, no injection seam - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: the permission check is a classmethod the handler calls directly, no injection seam @@ -7603,7 +7603,7 @@ async def test_get_and_validate_existing_key(): ) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", return_value="hashed-test-key-123", ): result = await _get_and_validate_existing_key( @@ -7622,7 +7622,7 @@ async def test_get_and_validate_existing_key(): ) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", return_value="hashed-non-existent-key", ): with pytest.raises(ProxyException) as exc_info: @@ -7704,7 +7704,7 @@ async def test_process_single_key_update(): # Mock _delete_cache_key_object with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache: mock_delete_cache.return_value = None @@ -7714,7 +7714,7 @@ async def test_process_single_key_update(): # Mock _hash_token_if_needed with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", return_value="hashed-test-key-123", ): # Mock KeyManagementEventHooks @@ -7853,7 +7853,7 @@ async def test_bulk_update_keys_success(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint" ): with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ): with patch("litellm.proxy._types.hash_token") as mock_hash: mock_hash.side_effect = ["hashed-key-1", "hashed-key-2"] @@ -7865,7 +7865,7 @@ async def test_bulk_update_keys_success(monkeypatch): }[token] with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", side_effect=_hash_for_bulk_success, ): with patch( @@ -7981,7 +7981,7 @@ async def test_bulk_update_keys_partial_failures(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint" ): with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ): with patch("litellm.proxy._types.hash_token") as mock_hash: mock_hash.return_value = "hashed-key-1" @@ -7993,7 +7993,7 @@ async def test_bulk_update_keys_partial_failures(monkeypatch): }[token] with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed", + "litellm.proxy.management_endpoints.key_management_endpoints.hash_token_if_needed", side_effect=_hash_for_bulk_partial, ): with patch( @@ -8186,7 +8186,7 @@ async def test_reset_key_spend_success(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" ) as mock_check_admin, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_hash_token.return_value = hashed_key @@ -8306,7 +8306,7 @@ async def test_reset_key_spend_resets_budget_windows(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" ) as mock_check_admin, patch( # test-quality-ok: no HTTP boundary; same pattern as test_reset_key_spend_success - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_hash_token.return_value = hashed_key @@ -8415,7 +8415,7 @@ async def test_reset_key_spend_no_budget_limits_skips_window_reset(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" ) as mock_check_admin, patch( # test-quality-ok: no HTTP boundary; same pattern as test_reset_key_spend_success - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_hash_token.return_value = hashed_key @@ -8461,7 +8461,7 @@ async def test_delete_cache_key_object_broadcasts_invalidation(monkeypatch): "litellm.proxy.auth.auth_checks.publish_auth_cache_invalidation" ) as mock_publish: mock_publish.return_value = None - await _delete_cache_key_object( + await delete_cache_key_object( hashed_token="hashed-broadcast-key", user_api_key_cache=real_user_api_key_cache, proxy_logging_obj=mock_proxy_logging_obj, @@ -8519,7 +8519,7 @@ async def test_update_key_spend_updates_counter(monkeypatch): ) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache: mock_delete_cache.return_value = None @@ -8616,7 +8616,7 @@ async def test_reset_key_spend_success_team_admin(monkeypatch): with ( patch("litellm.proxy.proxy_server.hash_token") as mock_hash_token, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_hash_token.return_value = hashed_key @@ -8830,7 +8830,7 @@ async def test_reset_key_spend_hashed_key(monkeypatch): "litellm.proxy.management_endpoints.key_management_endpoints._check_proxy_or_team_admin_for_key" ) as mock_check_admin, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object" ) as mock_delete_cache, ): mock_check_admin.return_value = None @@ -9326,7 +9326,7 @@ async def test_rotate_master_key_reencrypts_model_params_in_place( from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) # Setup mock prisma client @@ -9394,7 +9394,7 @@ async def test_rotate_master_key_reencrypts_model_params_in_place( "litellm.proxy.proxy_server.proxy_config", mock_proxy_config, ): - await _rotate_master_key( + await rotate_master_key( prisma_client=mock_prisma_client, user_api_key_dict=user_api_key_dict, current_master_key="sk-old-master-key", @@ -11016,7 +11016,7 @@ def _setup_block_unblock_mocks(monkeypatch, mock_key_team_id=None): pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) @@ -11270,7 +11270,7 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) @@ -11431,7 +11431,7 @@ async def test_update_key_throttle_unchanged_allows_non_budget_edit_for_internal pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", _noop, ) monkeypatch.setattr( @@ -11669,7 +11669,7 @@ async def test_update_key_team_member_with_permission_can_update_non_budget( mock_enforce_unique_key_alias, ) monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", mock_delete_cache_key_object, ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -12811,11 +12811,11 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_duration(monk new_callable=AsyncMock, ), patch( # test-quality-ok: archival path is outside upperbound rejection - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", new_callable=AsyncMock, ) as persist_deleted_verification_tokens, patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), ): @@ -12861,11 +12861,11 @@ async def test_execute_virtual_key_regeneration_changed_alias_must_match_key_ali new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), ): @@ -12921,7 +12921,7 @@ async def test_execute_virtual_key_regeneration_allows_within_limit_duration(mon new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -12967,11 +12967,11 @@ async def test_execute_virtual_key_regeneration_rejects_when_custom_key_update_h new_callable=AsyncMock, ) as insert_deprecated_key, patch( # test-quality-ok: archival path is outside policy rejection - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", new_callable=AsyncMock, ) as persist_deleted_verification_tokens, patch( # test-quality-ok: cache eviction is outside policy rejection - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: rotation callback is outside policy rejection @@ -13027,11 +13027,11 @@ async def test_execute_virtual_key_regeneration_allows_when_custom_key_update_ho new_callable=AsyncMock, ), patch( # test-quality-ok: verify archival follows policy approval - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", new_callable=AsyncMock, ) as persist_deleted_verification_tokens, patch( # test-quality-ok: cache eviction is outside policy approval - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: rotation callback is outside policy approval @@ -13080,7 +13080,7 @@ async def test_execute_virtual_key_regeneration_skips_custom_key_update_hook_wit new_callable=AsyncMock, ), patch( # test-quality-ok: cache eviction is outside unchanged request - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: rotation callback is outside unchanged request @@ -13129,7 +13129,7 @@ async def test_execute_virtual_key_regeneration_hides_the_untouched_modal_expiry new_callable=AsyncMock, ), patch( # test-quality-ok: cache eviction is outside the hook input - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: rotation callback is outside the hook input @@ -13197,13 +13197,13 @@ def _regenerate_policy_mocks(policy, insert_deprecated_key: AsyncMock, persist: ) stack.enter_context( patch( # test-quality-ok: archival write must not run on a denied regenerate - "litellm.proxy.management_endpoints.key_management_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.key_management_endpoints.persist_deleted_verification_tokens", persist, ) ) stack.enter_context( patch( # test-quality-ok: cache eviction is outside the policy path - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ) ) @@ -13340,7 +13340,7 @@ async def test_update_key_fn_runs_custom_key_policy_on_the_effective_row(monkeyp with ( patch( # test-quality-ok: cache eviction is outside the policy path - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch("litellm.proxy.proxy_server.user_custom_key_policy", policy), # test-quality-ok: inject policy hook @@ -13387,7 +13387,7 @@ async def test_update_key_fn_rejects_when_custom_key_policy_denies(monkeypatch): async def _process_single_key_update_under_policy(prisma_client: AsyncMock, data: UpdateKeyRequest, policy): with ( patch( # test-quality-ok: cache eviction is outside the policy path - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: update callback is outside the policy path @@ -13492,7 +13492,7 @@ def _setup_update_key_fn_object_permission_mocks(monkeypatch, allowed: bool) -> _record_object_permission_writes(mock_prisma_client, events) monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_key_policy", _recording_policy(events, allowed)) monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", AsyncMock() + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", AsyncMock() ) return mock_prisma_client, events @@ -13628,7 +13628,7 @@ async def test_bulk_update_keys_runs_custom_key_policy_per_key(monkeypatch): with ( patch( # test-quality-ok: cache eviction is outside the policy path - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: update callback is outside the policy path @@ -13996,7 +13996,7 @@ async def test_regenerate_evicts_jwt_key_mapping_cache_so_next_jwt_call_gets_new new_callable=AsyncMock, ), patch( # test-quality-ok: key-object eviction is separate from the mapping eviction under test - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( # test-quality-ok: background rotation hook is irrelevant to cache eviction @@ -14090,7 +14090,7 @@ async def test_execute_virtual_key_regeneration_rejects_over_limit_max_budget(mo new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), ): @@ -14146,7 +14146,7 @@ async def test_execute_virtual_key_regeneration_skips_none_values(monkeypatch): new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -14193,7 +14193,7 @@ async def test_execute_virtual_key_regeneration_no_upperbound_config_is_noop(mon new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -14685,7 +14685,7 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash(): return_value=None, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ) as mock_delete_cache, patch( @@ -14875,7 +14875,7 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ) as mock_delete_cache, patch( @@ -14968,7 +14968,7 @@ def _setup_team_keys_mocks( f"{_BULK_PKG}.prepare_key_update_data", AsyncMock(return_value={"max_budget": 50.0}), ) - monkeypatch.setattr(f"{_BULK_PKG}._delete_cache_key_object", AsyncMock()) + monkeypatch.setattr(f"{_BULK_PKG}.delete_cache_key_object", AsyncMock()) monkeypatch.setattr( f"{_BULK_PKG}.KeyManagementEventHooks.async_key_updated_hook", AsyncMock() ) @@ -14980,7 +14980,7 @@ def _setup_team_keys_mocks( ) if hash_identity: # Tests use already-hashed tokens; the raw-sk regression opts out. - monkeypatch.setattr(f"{_BULK_PKG}._hash_token_if_needed", lambda token: token) + monkeypatch.setattr(f"{_BULK_PKG}.hash_token_if_needed", lambda token: token) return mock_prisma @@ -15565,7 +15565,7 @@ def _patch_regenerate_side_effects(): new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -15672,7 +15672,7 @@ async def test_regenerate_premium_gate_allows_actual_master_key_holder(): patch("litellm.proxy.proxy_server.master_key", master), patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._rotate_master_key", + "litellm.proxy.management_endpoints.key_management_endpoints.rotate_master_key", new_callable=AsyncMock, ), ): @@ -17158,7 +17158,7 @@ async def _list_keys_capture_helper_kwargs(user_api_key_dict, **list_kwargs): return_value=mock_user_info, ): with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", helper, ): await list_keys( @@ -18467,7 +18467,7 @@ async def test_list_keys_forwards_expires_filter(expires_value, expected_forward return_value=mock_user_info, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", mock_helper, ), ): @@ -18506,7 +18506,7 @@ async def test_list_keys_without_expires_param_forwards_none(): return_value=mock_user_info, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + "litellm.proxy.management_endpoints.key_management_endpoints.list_key_helper", mock_helper, ), ): @@ -18546,7 +18546,7 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) mock_prisma_client = AsyncMock() @@ -18580,7 +18580,7 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( "litellm.proxy.proxy_server.proxy_config", mock_proxy_config, ): - await _rotate_master_key( + await rotate_master_key( prisma_client=mock_prisma_client, user_api_key_dict=user_api_key_dict, current_master_key="sk-old-master-key", @@ -18604,7 +18604,7 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch): encrypt_value_helper, ) from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) @@ -18639,7 +18639,7 @@ async def test_rotate_master_key_rotates_search_tools(monkeypatch): user_id="test-user", ) - await _rotate_master_key( + await rotate_master_key( prisma_client=mock_prisma_client, user_api_key_dict=user_api_key_dict, current_master_key="sk-old-master-key", @@ -18901,7 +18901,7 @@ def _wire_update_key_fn(monkeypatch, existing_key): pass monkeypatch.setattr( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", _noop, ) monkeypatch.setattr( @@ -19570,7 +19570,7 @@ async def test_update_key_syncs_access_group_assigned_key_ids_in_both_directions with ( patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -19657,7 +19657,7 @@ async def test_update_key_leaves_access_groups_alone_when_field_is_unset(monkeyp _setup_update_key_mocks(monkeypatch, mock_prisma_client) with patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ): await update_key_fn( @@ -19705,7 +19705,7 @@ async def test_bulk_update_keys_syncs_access_group_assigned_key_ids(monkeypatch) with ( patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -19882,7 +19882,7 @@ async def test_regenerate_key_repoints_access_group_assigned_key_ids(monkeypatch new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -19961,7 +19961,7 @@ async def test_key_write_paths_revoke_the_key_cache_before_syncing_access_groups with ( patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, side_effect=lambda **kwargs: order.append("revoke_key_cache"), ), @@ -20035,7 +20035,7 @@ async def test_update_key_syncs_many_access_groups_in_one_statement_per_directio with ( patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -20118,7 +20118,7 @@ async def test_regenerate_key_repoints_live_membership_not_the_key_row_it_read( new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -20899,7 +20899,7 @@ async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch new_callable=AsyncMock, ), patch( - "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + "litellm.proxy.management_endpoints.key_management_endpoints.delete_cache_key_object", new_callable=AsyncMock, ), patch( @@ -20967,7 +20967,7 @@ async def test_key_update_evicts_object_permission_before_key_object(monkeypatch """ from litellm.proxy._types import LiteLLM_ObjectPermissionBase from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key - from litellm.proxy.utils import _hash_token_if_needed + from litellm.proxy.utils import hash_token_if_needed deleted: list[str] = [] @@ -21028,7 +21028,7 @@ async def test_key_update_evicts_object_permission_before_key_object(monkeypatch ) assert deleted.index(object_permission_cache_key(permission_id)) < deleted.index( - _hash_token_if_needed("sk-lit5479") + hash_token_if_needed("sk-lit5479") ), deleted @@ -21207,7 +21207,7 @@ async def test_rotate_master_key_reencrypts_guardrail_params(monkeypatch): ) from litellm.proxy.management_endpoints import key_management_endpoints from litellm.proxy.management_endpoints.key_management_endpoints import ( - _rotate_master_key, + rotate_master_key, ) for rotator in ( @@ -21232,7 +21232,7 @@ async def test_rotate_master_key_reencrypts_guardrail_params(monkeypatch): mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[guardrail_row]) mock_prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=1) - await _rotate_master_key( + await rotate_master_key( prisma_client=mock_prisma_client, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"), current_master_key="sk-old-master-key", diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index 3c2e099e06c..21ff5534e03 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -422,7 +422,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -669,7 +669,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -788,7 +788,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), patch( @@ -924,7 +924,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -971,7 +971,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1028,7 +1028,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( # test-quality-ok: endpoint test must patch module globals - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1105,7 +1105,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1149,7 +1149,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1195,7 +1195,7 @@ class TestListMCPServers: AsyncMock(return_value=mock_health_result), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1237,7 +1237,7 @@ class TestListMCPServers: side_effect=lambda sid: config_server if sid == "serper_custom_dev" else None ) mock_manager.get_mcp_server_by_name = MagicMock(return_value=None) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record( server_id="serper_custom_dev", alias="Serper MCP", @@ -1264,7 +1264,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1281,7 +1281,7 @@ class TestListMCPServers: assert result.server_id == "serper_custom_dev" assert result.status == "healthy" mock_manager.get_mcp_server_by_id.assert_called_with("serper_custom_dev") - mock_manager._build_mcp_server_table.assert_called_once() + mock_manager.build_mcp_server_table.assert_called_once() @pytest.mark.asyncio async def test_fetch_single_mcp_server_from_registry_by_name_passes_client_ip(self): @@ -1299,7 +1299,7 @@ class TestListMCPServers: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id = MagicMock(return_value=None) mock_manager.get_mcp_server_by_name = MagicMock(return_value=config_server) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record( server_id="serper_custom_dev", alias="Serper MCP", @@ -1328,7 +1328,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), ): @@ -1363,7 +1363,7 @@ class TestListMCPServers: side_effect=lambda sid: config_server if sid == "restricted_server" else None ) mock_manager.get_mcp_server_by_name = MagicMock(return_value=None) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record( server_id="restricted_server", alias="Restricted MCP", @@ -1391,7 +1391,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), ): @@ -1431,7 +1431,7 @@ class TestListMCPServers: side_effect=lambda sid: config_server if sid == "allowed_config_server" else None ) mock_manager.get_mcp_server_by_name = MagicMock(return_value=None) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record( server_id="allowed_config_server", alias="Allowed MCP", @@ -1458,7 +1458,7 @@ class TestListMCPServers: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), ): @@ -1545,7 +1545,7 @@ class TestListMCPServers: AsyncMock(return_value=[generate_mock_mcp_server_db_record(server_id="env-server")]), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), ): @@ -1707,7 +1707,7 @@ class TestTeamScopedMCPServerAccess: with ( patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), patch( @@ -1743,13 +1743,13 @@ class TestTeamScopedMCPServerAccess: mock_server = generate_mock_mcp_server_config_record(server_id="server-1", name="Team Server") mock_manager = MagicMock() mock_manager.get_mcp_server_by_id = MagicMock(return_value=mock_server) - mock_manager._build_mcp_server_table = MagicMock( + mock_manager.build_mcp_server_table = MagicMock( return_value=generate_mock_mcp_server_db_record(server_id="server-1") ) with ( patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=False, ), patch( @@ -1783,7 +1783,7 @@ class TestTeamScopedMCPServerAccess: with ( patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -1836,7 +1836,7 @@ class TestFetchAllMCPServersOrdering: mock_manager, ), patch( # test-quality-ok: admin view is derived from module-global proxy settings - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( # test-quality-ok: auth contexts need a live prisma client @@ -1902,10 +1902,10 @@ class TestTemporaryMCPSessionEndpoints: mock_manager, ): from litellm.proxy.management_endpoints.mcp_management_endpoints import ( - _inherit_credentials_from_existing_server, + inherit_credentials_from_existing_server, ) - updated_payload = _inherit_credentials_from_existing_server(payload) + updated_payload = inherit_credentials_from_existing_server(payload) assert updated_payload.credentials == { "auth_value": "token-abc", @@ -1946,10 +1946,10 @@ class TestTemporaryMCPSessionEndpoints: mock_manager, ): from litellm.proxy.management_endpoints.mcp_management_endpoints import ( - _inherit_credentials_from_existing_server, + inherit_credentials_from_existing_server, ) - return _inherit_credentials_from_existing_server(payload) + return inherit_credentials_from_existing_server(payload) @pytest.mark.parametrize( "credentials", @@ -2509,7 +2509,7 @@ class TestTemporaryMCPSessionEndpoints: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id.return_value = registry_server mock_manager.get_mcp_server_by_name.return_value = None - mock_manager._build_mcp_server_table.return_value = generate_mock_mcp_server_db_record(server_id="server-x") + mock_manager.build_mcp_server_table.return_value = generate_mock_mcp_server_db_record(server_id="server-x") mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server-x"]) with ( @@ -2555,7 +2555,7 @@ class TestTemporaryMCPSessionEndpoints: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id.return_value = registry_server mock_manager.get_mcp_server_by_name.return_value = None - mock_manager._build_mcp_server_table.return_value = generate_mock_mcp_server_db_record(server_id="server-x") + mock_manager.build_mcp_server_table.return_value = generate_mock_mcp_server_db_record(server_id="server-x") def allowed_for(auth): return ["server-x"] if auth.team_id == "team-with-mcp-grant" else [] @@ -2750,11 +2750,11 @@ class TestTemporaryMCPSessionEndpoints: with ( patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_auth_builder", AsyncMock(return_value=expected_auth), ) as auth_builder_mock, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value={}), ), patch( @@ -2785,11 +2785,11 @@ class TestTemporaryMCPSessionEndpoints: with ( patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_auth_builder", AsyncMock(return_value=expected_auth), ) as auth_builder_mock, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value={}), ), patch( @@ -2834,11 +2834,11 @@ class TestTemporaryMCPSessionEndpoints: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_auth_builder", AsyncMock(return_value=expected_auth), ) as auth_builder_mock, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value={}), ), patch( @@ -2888,11 +2888,11 @@ class TestTemporaryMCPSessionEndpoints: mock_manager, ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_auth_builder", AsyncMock(return_value=expected_auth), ) as auth_builder_mock, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value={}), ), patch( @@ -3726,7 +3726,7 @@ class TestTemporaryMCPSessionEndpoints: return_value=nullcontext(server), ) as get_server, patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value=request_body), ) as read_body, patch( @@ -3790,7 +3790,7 @@ class TestTemporaryMCPSessionEndpoints: return_value=nullcontext(server), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value=request_body), ), patch( @@ -3835,7 +3835,7 @@ class TestTemporaryMCPSessionEndpoints: return_value=nullcontext(server), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + "litellm.proxy.management_endpoints.mcp_management_endpoints.read_request_body", AsyncMock(return_value=request_body), ), patch( @@ -5444,7 +5444,7 @@ async def test_store_mcp_oauth_user_credential_returns_status(): new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -5513,7 +5513,7 @@ async def test_store_mcp_oauth_user_credential_blocked_when_identity_binding_enf new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)), ), patch( # test-quality-ok: mirrors the existing store-credential tests in this file - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch.object( # test-quality-ok: registry is a module-level singleton; injecting it would change the endpoint signature @@ -5607,7 +5607,7 @@ async def test_store_mcp_oauth_user_credential_invalidates_cached_token(): new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)), ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.mcp_management_endpoints.user_api_key_has_admin_view", return_value=True, ), patch( @@ -7360,7 +7360,7 @@ class TestPerUserCredentialConfigServerResolution: manager.get_mcp_server_by_id = MagicMock( side_effect=lambda sid: config_server if sid == self.CONFIG_SERVER_ID else None ) - manager._build_mcp_server_table = MagicMock(return_value=record) + manager.build_mcp_server_table = MagicMock(return_value=record) manager.get_allowed_mcp_servers = AsyncMock(return_value=[]) return manager @@ -7488,7 +7488,7 @@ class TestPerUserCredentialConfigServerResolution: manager.get_mcp_server_by_id = MagicMock( return_value=generate_mock_mcp_server_config_record(server_id=self.CONFIG_SERVER_ID) ) - manager._build_mcp_server_table = MagicMock(return_value=env_var_server) + manager.build_mcp_server_table = MagicMock(return_value=env_var_server) manager.get_allowed_mcp_servers = AsyncMock(return_value=[self.CONFIG_SERVER_ID]) merge_mock = AsyncMock(return_value={"CORP_USERNAME": "alice"}) with ( @@ -8700,7 +8700,7 @@ class TestPinMCPServerTools: } ) ) - manager._get_tools_from_server = AsyncMock( + manager.get_tools_from_server = AsyncMock( return_value=[ MCPTool(name=name, description=description, inputSchema=schema) for name, description, schema in upstream_tools @@ -8747,7 +8747,7 @@ class TestPinMCPServerTools: "count_notes": PinnedMCPTool(description="", input_schema={}), } assert result == expected - listing = manager._get_tools_from_server.await_args.kwargs + listing = manager.get_tools_from_server.await_args.kwargs assert listing["server"].pinned_tools is None assert listing["server"].tool_name_to_description is None assert listing["proxy_logging_obj"] is None @@ -8778,7 +8778,7 @@ class TestPinMCPServerTools: assert result == {"server_id": "srv-1", "status": "unpinned"} assert store_mock.await_args.args[1:] == ("srv-1", None) assert store_mock.await_args.kwargs == {"touched_by": "admin"} - manager._get_tools_from_server.assert_not_awaited() + manager.get_tools_from_server.assert_not_awaited() manager.reload_servers_from_database.assert_awaited_once() @pytest.mark.asyncio @@ -8822,7 +8822,7 @@ class TestPinMCPServerTools: assert (pin_exc.value.status_code, unpin_exc.value.status_code) == (403, 403) store_mock.assert_not_awaited() - manager._get_tools_from_server.assert_not_awaited() + manager.get_tools_from_server.assert_not_awaited() @pytest.mark.asyncio async def test_pin_unknown_server_is_404(self): 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 4439a3511df..f44b2b57ccd 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -176,7 +176,7 @@ class MockProxyConfig: self.success = success self.deployment_called = False - async def _add_deployment_locked(self, prisma_client, proxy_logging_obj): + async def add_deployment_locked(self, prisma_client, proxy_logging_obj): self.deployment_called = True if not self.success: raise Exception("Failed to add deployment") @@ -830,7 +830,7 @@ class TestClearCache: mock_router.model_list = ["openai/gpt-4o", "openai/gpt-4o-mini"] mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -888,7 +888,7 @@ class TestClearCache: mock_router.complexity_routers = {"db-complexity-router": MagicMock(), "config-router": MagicMock()} mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -922,7 +922,7 @@ class TestClearCache: assert "config-router" in mock_router.complexity_routers # Should have called the already-locked reload to restore DB models - mock_config._add_deployment_locked.assert_called_once_with( + mock_config.add_deployment_locked.assert_called_once_with( prisma_client=mock_prisma, proxy_logging_obj=mock_logging ) @@ -967,7 +967,7 @@ class TestClearCache: mock_router.quality_routers = {} mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -1023,7 +1023,7 @@ class TestClearCachePreservesConfigRouters: } mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -1065,7 +1065,7 @@ class TestClearCachePreservesConfigRouters: mock_router.complexity_routers = {"shared-name": MagicMock()} mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -1109,7 +1109,7 @@ class TestClearCachePreservesConfigRouters: mock_router.adaptive_routers = {"a1": MagicMock()} mock_config = MagicMock() - mock_config._add_deployment_locked = AsyncMock( + mock_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -1632,7 +1632,7 @@ class TestTeamModelSiblingRouting: team_model_add to register the public name on the team's models list. """ from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_team_model_to_db, + add_team_model_to_db, ) from litellm.types.router import ModelInfo @@ -1659,7 +1659,7 @@ class TestTeamModelSiblingRouting: ) with ( patch( - "litellm.proxy.management_endpoints.model_management_endpoints._add_model_to_db", + "litellm.proxy.management_endpoints.model_management_endpoints.add_model_to_db", side_effect=mock_add_model_to_db, ), patch( @@ -1667,7 +1667,7 @@ class TestTeamModelSiblingRouting: mock_team_model_add, ), ): - await _add_team_model_to_db( + await add_team_model_to_db( model_params=dep, user_api_key_dict=user, prisma_client=prisma_client, @@ -2521,7 +2521,7 @@ class TestAddAndDeleteModelLifecycle: mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_proxy_config = MagicMock() - mock_proxy_config._add_deployment_locked = AsyncMock( + mock_proxy_config.add_deployment_locked = AsyncMock( return_value=ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) ) @@ -2644,7 +2644,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()) as mock_refresh, + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()) as mock_refresh, ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -2720,7 +2720,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -2793,7 +2793,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()) as mock_refresh, + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()) as mock_refresh, ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=deleted_id), @@ -2872,7 +2872,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", mock_router), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()) as mock_refresh, + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()) as mock_refresh, ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -2948,7 +2948,7 @@ class TestDeleteTeamBYOKModelGhost: patch(f"{_PS}.llm_router", mock_router), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -3023,7 +3023,7 @@ class TestDeleteModelTeamAuth: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id=model_id), @@ -3059,7 +3059,7 @@ class TestDeleteModelTeamAuth: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): with pytest.raises(ProxyException) as exc_info: await delete_model_endpoint( @@ -3124,7 +3124,7 @@ class TestDeleteModelTeamAuth: patch(f"{_PS}.llm_router", MagicMock()), patch(f"{_PS}.proxy_logging_obj", MagicMock()), patch(f"{_PS}.user_api_key_cache", MagicMock()), - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), ): with pytest.raises(ProxyException) as exc_info: await delete_model_endpoint( @@ -5201,7 +5201,7 @@ class TestConcurrentModelWritesDoNotEvictEachOther: depth -= 1 return ReconcileOutcome(still_desired=frozenset(), live_after=frozenset()) - monkeypatch.setattr(ProxyConfig, "_add_deployment_locked", fake_locked) + monkeypatch.setattr(ProxyConfig, "add_deployment_locked", fake_locked) config = ProxyConfig() await asyncio.gather( @@ -5244,7 +5244,7 @@ class TestConcurrentModelWritesDoNotEvictEachOther: async def fake_locked(self, **kwargs): return ReconcileOutcome(still_desired=frozenset({"m-db"}), live_after=frozenset({"m-db"})) - monkeypatch.setattr(ProxyConfig, "_add_deployment_locked", fake_locked) + monkeypatch.setattr(ProxyConfig, "add_deployment_locked", fake_locked) outcome = await asyncio.wait_for(clear_cache(), timeout=5) @@ -6408,7 +6408,7 @@ class TestStrategyRouterWriteValidation: lock holder waiting for a connection the waiters are occupying.""" from contextlib import asynccontextmanager - from litellm.proxy.management_endpoints.model_management_endpoints import _add_team_model_to_db + from litellm.proxy.management_endpoints.model_management_endpoints import add_team_model_to_db from litellm.types.router import ModelInfo events: list[str] = [] @@ -6438,7 +6438,7 @@ class TestStrategyRouterWriteValidation: side_effect=team_model_add, ), ): - result = await _add_team_model_to_db( + result = await add_team_model_to_db( model_params=deployment, user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), prisma_client=MagicMock(), @@ -8341,7 +8341,7 @@ class TestAddModelToDbBlocked: @pytest.mark.asyncio async def test_add_model_to_db_writes_blocked_true(self): from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, + add_model_to_db, ) mock_prisma = MagicMock() @@ -8351,7 +8351,7 @@ class TestAddModelToDbBlocked: with patch( # test-quality-ok: the proxy wiring under test is what this patches "litellm.proxy.proxy_server.master_key", "sk-test-master" ): # test-quality-ok: the proxy wiring under test is what this patches - await _add_model_to_db( + await add_model_to_db( model_params=self._deployment(True), user_api_key_dict=admin, prisma_client=mock_prisma ) @@ -8361,7 +8361,7 @@ class TestAddModelToDbBlocked: @pytest.mark.asyncio async def test_add_model_to_db_writes_blocked_false(self): from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, + add_model_to_db, ) mock_prisma = MagicMock() @@ -8371,7 +8371,7 @@ class TestAddModelToDbBlocked: with patch( # test-quality-ok: the proxy wiring under test is what this patches "litellm.proxy.proxy_server.master_key", "sk-test-master" ): # test-quality-ok: the proxy wiring under test is what this patches - await _add_model_to_db( + await add_model_to_db( model_params=self._deployment(False), user_api_key_dict=admin, prisma_client=mock_prisma ) @@ -8383,7 +8383,7 @@ class TestAddModelToDbBlocked: """None means "don't set it" -- the Prisma column defaults to False -- not "explicitly unblocked", so the key must be absent from the write entirely.""" from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, + add_model_to_db, ) mock_prisma = MagicMock() @@ -8393,7 +8393,7 @@ class TestAddModelToDbBlocked: with patch( # test-quality-ok: the proxy wiring under test is what this patches "litellm.proxy.proxy_server.master_key", "sk-test-master" ): # test-quality-ok: the proxy wiring under test is what this patches - await _add_model_to_db( + await add_model_to_db( model_params=self._deployment(None), user_api_key_dict=admin, prisma_client=mock_prisma ) @@ -8944,7 +8944,7 @@ class TestOneCredentialFeedsManyModelsNoWifCopy: async def test_two_discovered_models_share_the_credential_reference_only(self): from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, + add_model_to_db, ) from litellm.types.router import ModelInfo @@ -8957,7 +8957,7 @@ class TestOneCredentialFeedsManyModelsNoWifCopy: "litellm.proxy.proxy_server.master_key", "sk-test-master" ), patch( # test-quality-ok: the proxy wiring under test is what this patches - "litellm.proxy.common_utils.encrypt_decrypt_utils._get_salt_key", return_value="sk-test-master" + "litellm.proxy.common_utils.encrypt_decrypt_utils.get_salt_key", return_value="sk-test-master" ), ): for i, discovered_id in enumerate(["claude-a", "claude-b"]): @@ -8969,7 +8969,7 @@ class TestOneCredentialFeedsManyModelsNoWifCopy: model_info=ModelInfo(id=f"dep-shared-{i}"), blocked=False, ) - await _add_model_to_db(model_params=model_params, user_api_key_dict=admin, prisma_client=mock_prisma) + await add_model_to_db(model_params=model_params, user_api_key_dict=admin, prisma_client=mock_prisma) assert mock_prisma.db.litellm_proxymodeltable.create.await_count == 2 for call in mock_prisma.db.litellm_proxymodeltable.create.await_args_list: @@ -9335,7 +9335,7 @@ class TestFederationGateScopesToWhatTheWriteTouches: patch(f"{_PS}.llm_router", MagicMock()), # test-quality-ok: proxy wiring under test patch(f"{_PS}.proxy_logging_obj", MagicMock()), # test-quality-ok: proxy wiring under test patch(f"{_PS}.user_api_key_cache", MagicMock()), # test-quality-ok: proxy wiring under test - patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), # test-quality-ok: proxy wiring under test + patch(f"{_MOD}.refresh_cached_team", new=AsyncMock()), # test-quality-ok: proxy wiring under test ): result = await delete_model_endpoint( model_info=ModelInfoDelete(id="m1"), diff --git a/tests/unit/proxy/management_endpoints/test_org_admin_team_access.py b/tests/unit/proxy/management_endpoints/test_org_admin_team_access.py index aab67dccf1d..4a4aaf13154 100644 --- a/tests/unit/proxy/management_endpoints/test_org_admin_team_access.py +++ b/tests/unit/proxy/management_endpoints/test_org_admin_team_access.py @@ -41,9 +41,7 @@ def _make_team(team_id="team-1", organization_id="org-1") -> LiteLLM_TeamTable: ) -def _make_user_key( - user_id="org-admin-user", role=LitellmUserRoles.INTERNAL_USER.value -) -> UserAPIKeyAuth: +def _make_user_key(user_id="org-admin-user", role=LitellmUserRoles.INTERNAL_USER.value) -> UserAPIKeyAuth: return UserAPIKeyAuth(user_id=user_id, user_role=role) @@ -57,9 +55,7 @@ def _make_membership(user_id, org_id, role="org_admin"): ) -def _make_caller_user( - user_id="org-admin-user", org_id="org-1", org_role="org_admin" -) -> LiteLLM_UserTable: +def _make_caller_user(user_id="org-admin-user", org_id="org-1", org_role="org_admin") -> LiteLLM_UserTable: return LiteLLM_UserTable( user_id=user_id, organization_memberships=[_make_membership(user_id, org_id, org_role)], @@ -76,9 +72,7 @@ def _patch_org_admin_deps(get_user_return): ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock(), create=True), patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(), create=True), - patch( - "litellm.proxy.proxy_server.user_api_key_cache", MagicMock(), create=True - ), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock(), create=True), ) @@ -133,9 +127,7 @@ class TestValidateMembership: team = _make_team(organization_id="org-1") key = _make_user_key(user_id="random-user") - caller = _make_caller_user( - user_id="random-user", org_id="org-2", org_role="user" - ) + caller = _make_caller_user(user_id="random-user", org_id="org-2", org_role="user") p1, p2, p3, p4 = _patch_org_admin_deps(caller) with p1, p2, p3, p4: @@ -150,9 +142,7 @@ class TestValidateMembership: ) team = _make_team(team_id="team-1") - key = UserAPIKeyAuth( - team_id="team-1", user_role=LitellmUserRoles.INTERNAL_USER.value - ) + key = UserAPIKeyAuth(team_id="team-1", user_role=LitellmUserRoles.INTERNAL_USER.value) await validate_membership(user_api_key_dict=key, team_table=team) @@ -168,55 +158,49 @@ class TestUserIsOrgAdminRouteCheck: """ def test_no_candidate_org_ids_returns_false(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin user = LiteLLM_UserTable( user_id="org-admin-user", organization_memberships=[_make_membership("org-admin-user", "org-1")], ) - result = _user_is_org_admin(request_data={}, user_object=user) + result = user_is_org_admin(request_data={}, user_object=user) assert result is False, "Must NOT grant blanket access when no org in request" def test_matching_org_id_returns_true(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin user = LiteLLM_UserTable( user_id="org-admin-user", organization_memberships=[_make_membership("org-admin-user", "org-1")], ) - result = _user_is_org_admin( - request_data={"organization_id": "org-1"}, user_object=user - ) + result = user_is_org_admin(request_data={"organization_id": "org-1"}, user_object=user) assert result is True def test_non_matching_org_id_returns_false(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin user = LiteLLM_UserTable( user_id="org-admin-user", organization_memberships=[_make_membership("org-admin-user", "org-1")], ) - result = _user_is_org_admin( - request_data={"organization_id": "org-99"}, user_object=user - ) + result = user_is_org_admin(request_data={"organization_id": "org-99"}, user_object=user) assert result is False def test_organizations_list_field(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin user = LiteLLM_UserTable( user_id="org-admin-user", organization_memberships=[_make_membership("org-admin-user", "org-1")], ) - result = _user_is_org_admin( - request_data={"organizations": ["org-1"]}, user_object=user - ) + result = user_is_org_admin(request_data={"organizations": ["org-1"]}, user_object=user) assert result is True def test_none_user_object_returns_false(self): - from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin + from litellm.proxy.auth.auth_checks_organization import user_is_org_admin - result = _user_is_org_admin(request_data={}, user_object=None) + result = user_is_org_admin(request_data={}, user_object=None) assert result is False def test_user_list_in_self_managed_routes(self): diff --git a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py index 7d3bd049a03..a98f1344279 100644 --- a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py @@ -108,7 +108,7 @@ async def test_get_organization_daily_activity_admin_param_passing(monkeypatch): # Admin view -> skip membership restriction monkeypatch.setattr( - "litellm.proxy.management_endpoints.organization_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.organization_endpoints.user_api_key_has_admin_view", lambda _: True, ) @@ -175,7 +175,7 @@ async def test_get_organization_daily_activity_non_admin_defaults_to_admin_orgs( # Non-admin view monkeypatch.setattr( - "litellm.proxy.management_endpoints.organization_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.organization_endpoints.user_api_key_has_admin_view", lambda _: False, ) @@ -227,7 +227,7 @@ async def test_get_organization_daily_activity_non_admin_unauthorized_org_raises # Non-admin view monkeypatch.setattr( - "litellm.proxy.management_endpoints.organization_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.organization_endpoints.user_api_key_has_admin_view", lambda _: False, ) @@ -853,7 +853,7 @@ async def _run_update_organization_v2( mock_prisma_client.call_order = call_order monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) + monkeypatch.setattr(organization_endpoints, "verify_org_access", AsyncMock()) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1") await update_organization_v2( @@ -1019,7 +1019,7 @@ async def test_v2_rejects_caller_without_org_access(monkeypatch): mock_prisma_client = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr(organization_endpoints, "_user_has_admin_view", lambda _: False) + monkeypatch.setattr(organization_endpoints, "user_api_key_has_admin_view", lambda _: False) caller = MagicMock() caller.organization_memberships = [] @@ -1110,7 +1110,7 @@ async def test_v2_rejects_empty_object_permission(monkeypatch): mock_prisma_client = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) + monkeypatch.setattr(organization_endpoints, "verify_org_access", AsyncMock()) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1") with pytest.raises(HTTPException) as exc: @@ -1178,7 +1178,7 @@ async def _run_legacy_update_organization( mock_prisma_client.db.litellm_budgettable.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) + monkeypatch.setattr(organization_endpoints, "verify_org_access", AsyncMock()) request = MagicMock() request.json = AsyncMock(return_value=body) @@ -1299,7 +1299,7 @@ async def test_get_organization_daily_activity_non_admin_without_org_admin_role_ monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr( - "litellm.proxy.management_endpoints.organization_endpoints._user_has_admin_view", + "litellm.proxy.management_endpoints.organization_endpoints.user_api_key_has_admin_view", lambda _: False, ) diff --git a/tests/unit/proxy/management_endpoints/test_policy_endpoints.py b/tests/unit/proxy/management_endpoints/test_policy_endpoints.py index f5c4e9b69ff..e334305dd98 100644 --- a/tests/unit/proxy/management_endpoints/test_policy_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_policy_endpoints.py @@ -648,47 +648,47 @@ class TestApplyPoliciesDirectGuardrailNames: # Tests for competitor enrichment helper functions # --------------------------------------------------------------------------- from litellm.proxy.management_endpoints.policy_endpoints import ( - _build_all_names_per_competitor, - _build_comparison_blocked_words, - _build_competitor_guardrail_definitions, - _build_name_blocked_words, - _build_recommendation_blocked_words, - _build_refinement_prompt, - _clean_competitor_line, - _parse_variations_response, + build_all_names_per_competitor, + build_comparison_blocked_words, + build_competitor_guardrail_definitions, + build_name_blocked_words, + build_recommendation_blocked_words, + build_refinement_prompt, + clean_competitor_line, + parse_variations_response, ) class TestCleanCompetitorLine: - """Tests for _clean_competitor_line.""" + """Tests for clean_competitor_line.""" def test_strips_bullets_and_dashes(self): - assert _clean_competitor_line("- United Airlines") == "United Airlines" - assert _clean_competitor_line(" - JetBlue ") == "JetBlue" + assert clean_competitor_line("- United Airlines") == "United Airlines" + assert clean_competitor_line(" - JetBlue ") == "JetBlue" def test_strips_trailing_punctuation(self): - assert _clean_competitor_line("Delta Airlines.") == "Delta Airlines" - assert _clean_competitor_line("Southwest)") == "Southwest" + assert clean_competitor_line("Delta Airlines.") == "Delta Airlines" + assert clean_competitor_line("Southwest)") == "Southwest" def test_returns_none_for_empty(self): - assert _clean_competitor_line("") is None - assert _clean_competitor_line(" ") is None + assert clean_competitor_line("") is None + assert clean_competitor_line(" ") is None def test_returns_none_for_single_char(self): - assert _clean_competitor_line("A") is None - assert _clean_competitor_line(" - ") is None + assert clean_competitor_line("A") is None + assert clean_competitor_line(" - ") is None def test_plain_name(self): - assert _clean_competitor_line("Qatar Airways") == "Qatar Airways" + assert clean_competitor_line("Qatar Airways") == "Qatar Airways" class TestParseVariationsResponse: - """Tests for _parse_variations_response.""" + """Tests for parse_variations_response.""" def test_parses_standard_format(self): raw = "Delta Airlines: Delta Air Lines, DeltaAirlines, Delta\nUnited Airlines: United, UAL" competitors = ["Delta Airlines", "United Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) assert "Delta Airlines" in result assert "Delta Air Lines" in result["Delta Airlines"] assert "United" in result["United Airlines"] @@ -696,71 +696,71 @@ class TestParseVariationsResponse: def test_case_insensitive_matching(self): raw = "delta airlines: Delta Air Lines, DeltaAirlines" competitors = ["Delta Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) assert "Delta Airlines" in result assert len(result["Delta Airlines"]) == 2 def test_skips_lines_without_colon(self): raw = "This is a header\nDelta Airlines: Delta Air Lines" competitors = ["Delta Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) assert len(result) == 1 def test_skips_unknown_competitors(self): raw = "Unknown Corp: Foo, Bar\nDelta Airlines: Delta" competitors = ["Delta Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) assert "Unknown Corp" not in result assert "Delta Airlines" in result def test_filters_out_self_reference(self): raw = "Delta Airlines: Delta Airlines, Delta Air Lines" competitors = ["Delta Airlines"] - result = _parse_variations_response(raw, competitors) + result = parse_variations_response(raw, competitors) # "Delta Airlines" should be filtered out (same as canonical) assert "Delta Airlines" not in result["Delta Airlines"] assert "Delta Air Lines" in result["Delta Airlines"] def test_empty_input(self): - assert _parse_variations_response("", []) == {} + assert parse_variations_response("", []) == {} class TestBuildRefinementPrompt: - """Tests for _build_refinement_prompt.""" + """Tests for build_refinement_prompt.""" def test_includes_brand_name(self): - prompt = _build_refinement_prompt("add 10 more", ["Delta"], "Emirates") + prompt = build_refinement_prompt("add 10 more", ["Delta"], "Emirates") assert "Emirates" in prompt def test_includes_existing_competitors(self): - prompt = _build_refinement_prompt("add more", ["Delta", "United"], "Emirates") + prompt = build_refinement_prompt("add more", ["Delta", "United"], "Emirates") assert "Delta" in prompt assert "United" in prompt def test_includes_instruction(self): - prompt = _build_refinement_prompt("add 10 from Asia", ["Delta"], "Emirates") + prompt = build_refinement_prompt("add 10 from Asia", ["Delta"], "Emirates") assert "add 10 from Asia" in prompt def test_asks_for_new_names_only(self): - prompt = _build_refinement_prompt("add more", ["Delta"], "Emirates") + prompt = build_refinement_prompt("add more", ["Delta"], "Emirates") assert "NEW" in prompt class TestBuildAllNamesPerCompetitor: - """Tests for _build_all_names_per_competitor.""" + """Tests for build_all_names_per_competitor.""" def test_includes_canonical_and_variations(self): - result = _build_all_names_per_competitor( + result = build_all_names_per_competitor( ["Delta Airlines"], {"Delta Airlines": ["Delta", "DeltaAir"]} ) assert result["Delta Airlines"] == ["Delta Airlines", "Delta", "DeltaAir"] def test_no_variations(self): - result = _build_all_names_per_competitor(["Delta Airlines"], {}) + result = build_all_names_per_competitor(["Delta Airlines"], {}) assert result["Delta Airlines"] == ["Delta Airlines"] def test_multiple_competitors(self): - result = _build_all_names_per_competitor( + result = build_all_names_per_competitor( ["Delta", "United"], {"Delta": ["DL"], "United": ["UA"]}, ) @@ -770,11 +770,11 @@ class TestBuildAllNamesPerCompetitor: class TestBuildNameBlockedWords: - """Tests for _build_name_blocked_words.""" + """Tests for build_name_blocked_words.""" def test_basic_output(self): all_names = {"Delta": ["Delta", "DL"]} - result = _build_name_blocked_words(["Delta"], all_names) + result = build_name_blocked_words(["Delta"], all_names) keywords = [r["keyword"] for r in result] assert "Delta" in keywords assert "DL" in keywords @@ -782,18 +782,18 @@ class TestBuildNameBlockedWords: def test_descriptions_differ_for_variations(self): all_names = {"Delta": ["Delta", "DL"]} - result = _build_name_blocked_words(["Delta"], all_names) + result = build_name_blocked_words(["Delta"], all_names) descs = {r["keyword"]: r["description"] for r in result} assert "Competitor: Delta" == descs["Delta"] assert "variation" in descs["DL"].lower() class TestBuildRecommendationBlockedWords: - """Tests for _build_recommendation_blocked_words.""" + """Tests for build_recommendation_blocked_words.""" def test_generates_prefix_combinations(self): all_names = {"Delta": ["Delta"]} - result = _build_recommendation_blocked_words(["Delta"], all_names) + result = build_recommendation_blocked_words(["Delta"], all_names) keywords = [r["keyword"] for r in result] assert "try Delta" in keywords assert "use Delta" in keywords @@ -802,23 +802,23 @@ class TestBuildRecommendationBlockedWords: def test_includes_variations(self): all_names = {"Delta": ["Delta", "DL"]} - result = _build_recommendation_blocked_words(["Delta"], all_names) + result = build_recommendation_blocked_words(["Delta"], all_names) keywords = [r["keyword"] for r in result] assert "try DL" in keywords class TestBuildComparisonBlockedWords: - """Tests for _build_comparison_blocked_words.""" + """Tests for build_comparison_blocked_words.""" def test_generates_competitor_comparisons(self): all_names = {"Delta": ["Delta"]} - result = _build_comparison_blocked_words(["Delta"], all_names, "Emirates") + result = build_comparison_blocked_words(["Delta"], all_names, "Emirates") keywords = [r["keyword"] for r in result] assert "Delta is better" in keywords def test_generates_brand_comparisons_once(self): all_names = {"Delta": ["Delta"], "United": ["United"]} - result = _build_comparison_blocked_words( + result = build_comparison_blocked_words( ["Delta", "United"], all_names, "Emirates" ) keywords = [r["keyword"] for r in result] @@ -828,13 +828,13 @@ class TestBuildComparisonBlockedWords: def test_includes_variation_comparisons(self): all_names = {"Delta": ["Delta", "DL"]} - result = _build_comparison_blocked_words(["Delta"], all_names, "Emirates") + result = build_comparison_blocked_words(["Delta"], all_names, "Emirates") keywords = [r["keyword"] for r in result] assert "DL is better" in keywords class TestBuildCompetitorGuardrailDefinitions: - """Tests for _build_competitor_guardrail_definitions.""" + """Tests for build_competitor_guardrail_definitions.""" def test_populates_blocked_words_for_known_guardrail_names(self): definitions = [ @@ -847,7 +847,7 @@ class TestBuildCompetitorGuardrailDefinitions: "litellm_params": {"blocked_words": []}, }, ] - result = _build_competitor_guardrail_definitions( + result = build_competitor_guardrail_definitions( definitions, ["Delta"], "Emirates", {"Delta": ["DL"]} ) # Name blocker should have entries @@ -871,7 +871,7 @@ class TestBuildCompetitorGuardrailDefinitions: "litellm_params": {"blocked_words": ["original"]}, }, ] - result = _build_competitor_guardrail_definitions( + result = build_competitor_guardrail_definitions( definitions, ["Delta"], "Emirates" ) assert result[0]["litellm_params"]["blocked_words"] == ["original"] @@ -883,7 +883,7 @@ class TestBuildCompetitorGuardrailDefinitions: "litellm_params": {"blocked_words": []}, }, ] - _build_competitor_guardrail_definitions(definitions, ["Delta"], "Emirates") + build_competitor_guardrail_definitions(definitions, ["Delta"], "Emirates") # Original should be unchanged assert definitions[0]["litellm_params"]["blocked_words"] == [] @@ -898,7 +898,7 @@ class TestBuildCompetitorGuardrailDefinitions: "litellm_params": {"blocked_words": []}, }, ] - result = _build_competitor_guardrail_definitions( + result = build_competitor_guardrail_definitions( definitions, ["Delta"], "Emirates" ) for defn in result: diff --git a/tests/unit/proxy/management_endpoints/test_prompt_cache_prediction.py b/tests/unit/proxy/management_endpoints/test_prompt_cache_prediction.py index 987cacf7676..1ae978e9a2e 100644 --- a/tests/unit/proxy/management_endpoints/test_prompt_cache_prediction.py +++ b/tests/unit/proxy/management_endpoints/test_prompt_cache_prediction.py @@ -13,8 +13,8 @@ import litellm from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body, _safe_set_request_parsed_body -from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 +from litellm.proxy.common_utils.http_parsing_utils import read_request_body, safe_set_request_parsed_body +from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.llms.anthropic.prompt_cache_prediction import PromptPrefix, cache_scope, parse_prompt from litellm.proxy.hooks.prompt_cache_prediction import ( CacheObservation, @@ -224,7 +224,7 @@ def _app( if caller is not None: usage_cache: Final = InternalUsageCache(cache) configured_limiter: Final = ( - _PROXY_MaxParallelRequestsHandler_v3(usage_cache) if isinstance(limiter, str) else limiter + PROXY_MaxParallelRequestsHandler_v3(usage_cache) if isinstance(limiter, str) else limiter ) monkeypatch.setattr(proxy_server, "proxy_logging_obj", _ProxyLogging(usage_cache, configured_limiter)) app.dependency_overrides[endpoint.user_api_key_auth] = lambda: caller @@ -384,7 +384,7 @@ async def test_missing_or_unsupported_limiter_returns_unknown_before_counting( @pytest.mark.asyncio async def test_occupied_parallel_capacity_rejects_before_provider_count(monkeypatch: pytest.MonkeyPatch) -> None: cache: Final = DualCache() - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(InternalUsageCache(cache)) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3(InternalUsageCache(cache)) caller: Final = UserAPIKeyAuth(api_key=_CALLER, max_parallel_requests=1) app: Final = _app(monkeypatch, cache, caller=caller, counts=_unexpected_count, limiter=limiter) async with limiter.request_capacity(caller, "opus"): @@ -428,8 +428,8 @@ async def test_each_count_preserves_auth_cached_request_tag_limits( return await Counts()(model, api_key, body) async def authenticated_request(request: Request) -> UserAPIKeyAuth: - data: Final = await _read_request_body(request) - _safe_set_request_parsed_body(request, {**data, metadata_key: {"tags": ["cache-cost"]}}) + data: Final = await read_request_body(request) + safe_set_request_parsed_body(request, {**data, metadata_key: {"tags": ["cache-cost"]}}) return caller app: Final = _app(monkeypatch, DualCache(), caller=caller, counts=count) diff --git a/tests/unit/proxy/management_endpoints/test_ptu_model_settings.py b/tests/unit/proxy/management_endpoints/test_ptu_model_settings.py index d35b77f732c..5b4ce52baa5 100644 --- a/tests/unit/proxy/management_endpoints/test_ptu_model_settings.py +++ b/tests/unit/proxy/management_endpoints/test_ptu_model_settings.py @@ -16,7 +16,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token -from litellm.proxy.auth.auth_checks import _is_model_cost_zero +from litellm.proxy.auth.auth_checks import is_model_cost_zero from litellm.llms.gemini.cost_calculator import cost_per_web_search_request from litellm.proxy.management_endpoints.model_management_endpoints import ( _PTU_ZEROED_PRICING_FIELDS, @@ -677,8 +677,8 @@ class TestAddNewModelPtuGate: f"{endpoints}.ModelManagementAuthChecks.can_user_make_model_call", AsyncMock(return_value=True), ), - patch(f"{endpoints}._add_model_to_db", add_model_to_db), - patch(f"{endpoints}._add_team_model_to_db", add_team_model_to_db), + patch(f"{endpoints}.add_model_to_db", add_model_to_db), + patch(f"{endpoints}.add_team_model_to_db", add_team_model_to_db), ] @staticmethod @@ -1090,8 +1090,8 @@ class TestPtuDeploymentsAreNotBilledPerToken: ) ) router = Router(model_list=[priced.to_json(exclude_none=True)]) - assert _is_model_cost_zero(model="model_name_team-1_dep-ptu", llm_router=router) is False - assert _is_model_cost_zero(model="ptu-model", llm_router=router) is False + assert is_model_cost_zero(model="model_name_team-1_dep-ptu", llm_router=router) is False + assert is_model_cost_zero(model="ptu-model", llm_router=router) is False def test_an_unrelated_patch_heals_a_deployment_stored_before_this_rule(self): """Both blobs, because litellm_params wins over model_info wherever the two are merged.""" diff --git a/tests/unit/proxy/management_endpoints/test_session_endpoints.py b/tests/unit/proxy/management_endpoints/test_session_endpoints.py index d5960a88937..afde88c1c66 100644 --- a/tests/unit/proxy/management_endpoints/test_session_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_session_endpoints.py @@ -65,7 +65,7 @@ async def test_session_logout_revokes_presented_session(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", persist_mock, ), patch( @@ -100,7 +100,7 @@ async def test_session_logout_clears_token_cookie(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", AsyncMock(), ), patch( @@ -194,7 +194,7 @@ async def test_revoke_ui_session_keys_revokes_all_and_broadcasts(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", persist_mock, ), patch( @@ -228,7 +228,7 @@ async def test_revoke_ui_session_keys_keeps_callers_session(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", AsyncMock(), ), patch( @@ -275,7 +275,7 @@ async def test_revoke_ui_session_keys_failure_is_swallowed(): p2, p3, patch( - "litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens", + "litellm.proxy.management_endpoints.session_endpoints.persist_deleted_verification_tokens", AsyncMock(), ), ): diff --git a/tests/unit/proxy/management_endpoints/test_team_callback_endpoints.py b/tests/unit/proxy/management_endpoints/test_team_callback_endpoints.py index 331e3a7983c..f576bac4742 100644 --- a/tests/unit/proxy/management_endpoints/test_team_callback_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_team_callback_endpoints.py @@ -89,7 +89,7 @@ def stub_team_cache_refresh(): test_disable_team_logging_refreshes_cached_team. """ with patch( - "litellm.proxy.management_endpoints.team_callback_endpoints._refresh_cached_team", + "litellm.proxy.management_endpoints.team_callback_endpoints.refresh_cached_team", new_callable=AsyncMock, ) as refresh: yield refresh @@ -614,7 +614,7 @@ async def test_get_team_callbacks_decrypts_vars_stored_under_non_sensitive_keys( encrypted at rest under a key that later stops being masked on read. Without the decrypt step that value comes back as an unusable litellm_enc:: blob. """ - from litellm.proxy.common_utils.callback_utils import _CALLBACK_VAR_ENCRYPTED_PREFIX, is_sensitive_callback_key + from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX, is_sensitive_callback_key from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-32-bytes-aaaaaaaaaaaaaa") @@ -626,7 +626,7 @@ async def test_get_team_callbacks_decrypts_vars_stored_under_non_sensitive_keys( "callback_name": "langsmith", "callback_type": "success", "callback_vars": { - "langsmith_project": _CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("tenant-project"), + "langsmith_project": CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("tenant-project"), }, } ] @@ -655,11 +655,11 @@ async def test_get_team_callbacks_masks_values_that_fail_to_decrypt(monkeypatch) classified as sensitive it would otherwise reach the caller as an opaque blob that is indistinguishable from a real value. """ - from litellm.proxy.common_utils.callback_utils import _CALLBACK_VAR_ENCRYPTED_PREFIX + from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-32-bytes-aaaaaaaaaaaaaa") - stale = _CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("tenant-project") + stale = CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("tenant-project") monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-32-bytes-bbbbbbbbbbbbbb") metadata = { @@ -685,7 +685,7 @@ async def test_get_team_callbacks_masks_values_that_fail_to_decrypt(monkeypatch) assert response["data"]["success_callbacks"] == ["langsmith"] assert response["data"]["callback_vars"]["langsmith_project"] == "***REDACTED***" - assert _CALLBACK_VAR_ENCRYPTED_PREFIX not in json.dumps(response) + assert CALLBACK_VAR_ENCRYPTED_PREFIX not in json.dumps(response) @pytest.mark.asyncio @@ -752,7 +752,7 @@ async def test_disable_team_logging_stops_callbacks_registered_via_api(): the endpoint and then asks the real request-time resolver what the written row would do. """ - from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata + from litellm.proxy.litellm_pre_call_utils import get_dynamic_logging_metadata metadata = { "logging": [ @@ -780,7 +780,7 @@ async def test_disable_team_logging_stops_callbacks_registered_via_api(): written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"]) assert written["logging"] == [] - resolved = _get_dynamic_logging_metadata( + resolved = get_dynamic_logging_metadata( UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written), proxy_config=MagicMock(**{"load_team_config.return_value": {}}), ) @@ -859,7 +859,7 @@ async def test_add_team_callbacks_refreshes_cached_team(stub_team_cache_refresh) @pytest.mark.asyncio async def test_disable_team_logging_clears_both_metadata_shapes(): """A team carrying both shapes ends up with neither active.""" - from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata + from litellm.proxy.litellm_pre_call_utils import get_dynamic_logging_metadata metadata = { "logging": [ @@ -893,7 +893,7 @@ async def test_disable_team_logging_clears_both_metadata_shapes(): assert written["callback_settings"]["success_callback"] == [] assert written["callback_settings"]["failure_callback"] == [] - resolved = _get_dynamic_logging_metadata( + resolved = get_dynamic_logging_metadata( UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written), proxy_config=MagicMock(**{"load_team_config.return_value": {}}), ) @@ -1023,7 +1023,7 @@ async def test_delete_team_callback_leaves_the_other_callback_firing(): Asks the real request-time resolver what the written row would do, the same way the disable_logging regression test does. """ - from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata + from litellm.proxy.litellm_pre_call_utils import get_dynamic_logging_metadata mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata=_two_callback_metadata())) @@ -1040,7 +1040,7 @@ async def test_delete_team_callback_leaves_the_other_callback_firing(): ) written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"]) - resolved = _get_dynamic_logging_metadata( + resolved = get_dynamic_logging_metadata( UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written), proxy_config=MagicMock(**{"load_team_config.return_value": {}}), ) @@ -1219,7 +1219,7 @@ async def test_delete_team_callback_keeps_last_removal_from_reviving_legacy_shap dropping the key would fall through to a legacy callback_settings block and silently re-enable a destination the caller just removed. """ - from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata + from litellm.proxy.litellm_pre_call_utils import get_dynamic_logging_metadata metadata = { "logging": [ @@ -1253,7 +1253,7 @@ async def test_delete_team_callback_keeps_last_removal_from_reviving_legacy_shap assert written["logging"] == [] assert response.data.success_callbacks == () - resolved = _get_dynamic_logging_metadata( + resolved = get_dynamic_logging_metadata( UserAPIKeyAuth(api_key="hashed", team_id="team-1", team_metadata=written), proxy_config=MagicMock(**{"load_team_config.return_value": {}}), ) diff --git a/tests/unit/proxy/management_endpoints/test_team_endpoints.py b/tests/unit/proxy/management_endpoints/test_team_endpoints.py index 6c5447fce64..00f624ea518 100644 --- a/tests/unit/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_team_endpoints.py @@ -2313,7 +2313,7 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name): patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object", + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", new_callable=AsyncMock, ) as mock_cache_team, ): @@ -2492,7 +2492,7 @@ async def test_team_write_404s_when_row_vanishes_before_update(endpoint_name): patch("litellm.proxy.proxy_server.user_api_key_cache"), # test-quality-ok: proxy_server module global is the endpoint's only injection point patch("litellm.proxy.proxy_server.proxy_logging_obj"), # test-quality-ok: proxy_server module global is the endpoint's only injection point patch( # test-quality-ok: stubs the cache write so the test observes only the DB result handling - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object", + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", new_callable=AsyncMock, ), patch( # test-quality-ok: stubs the collaborator so the test pins the endpoint's own error contract @@ -2545,7 +2545,8 @@ async def test_update_team_team_member_budget_not_passed_to_db( patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object" + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", + new_callable=AsyncMock, ) as mock_cache_team, patch( "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" @@ -3113,7 +3114,8 @@ async def test_update_team_with_team_member_budget_duration( patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging, patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object" + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", + new_callable=AsyncMock, ) as mock_cache_team, patch( "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" @@ -11537,25 +11539,25 @@ def _non_admin_auth(): def test_check_passthrough_routes_caller_permission_team(): from litellm.proxy._types import NewTeamRequest from litellm.proxy.management_endpoints.common_utils import ( - _check_passthrough_routes_caller_permission, + check_passthrough_routes_caller_permission, ) admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) non_admin = _non_admin_auth() - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(allowed_passthrough_routes=["/foo/*"]), admin, entity="team" ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(), non_admin, entity="team" ) - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(allowed_passthrough_routes=[]), non_admin, entity="team" ) with pytest.raises(HTTPException) as exc: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(allowed_passthrough_routes=["/admin/*"]), non_admin, entity="team", @@ -11565,7 +11567,7 @@ def test_check_passthrough_routes_caller_permission_team(): assert "team" in str(exc.value.detail) with pytest.raises(HTTPException) as exc: - _check_passthrough_routes_caller_permission( + check_passthrough_routes_caller_permission( NewTeamRequest(metadata={"allowed_passthrough_routes": ["/admin/*"]}), non_admin, entity="team", @@ -11628,24 +11630,24 @@ async def test_update_team_blocks_non_admin_passthrough_routes(mock_db_client): def test_check_disable_global_guardrails_caller_permission_team(): from litellm.proxy._types import NewTeamRequest from litellm.proxy.management_endpoints.common_utils import ( - _check_disable_global_guardrails_caller_permission, + check_disable_global_guardrails_caller_permission, ) admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) non_admin = _non_admin_auth() - _check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, admin, entity="team") - _check_disable_global_guardrails_caller_permission(None, None, non_admin, entity="team") - _check_disable_global_guardrails_caller_permission(False, None, non_admin, entity="team") + check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, admin, entity="team") + check_disable_global_guardrails_caller_permission(None, None, non_admin, entity="team") + check_disable_global_guardrails_caller_permission(False, None, non_admin, entity="team") with pytest.raises(HTTPException) as exc: - _check_disable_global_guardrails_caller_permission(True, None, non_admin, entity="team") + check_disable_global_guardrails_caller_permission(True, None, non_admin, entity="team") assert exc.value.status_code == 403 assert "disable_global_guardrails" in str(exc.value.detail) assert "team" in str(exc.value.detail) with pytest.raises(HTTPException) as exc: - _check_disable_global_guardrails_caller_permission( + check_disable_global_guardrails_caller_permission( None, {"disable_global_guardrails": True}, non_admin, entity="team" ) assert exc.value.status_code == 403 @@ -12376,7 +12378,7 @@ async def _drive_team_write( _patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), _patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), _patch( - "litellm.proxy.management_endpoints.team_endpoints._refresh_cached_team", + "litellm.proxy.management_endpoints.team_endpoints.refresh_cached_team", new=AsyncMock(), ), ): @@ -13860,7 +13862,7 @@ async def test_team_member_update_role_change_emits_a_roster_audit_event(monkeyp AsyncMock(side_effect=[_team_info_as_read_from_db("user"), _team_info_as_read_from_db("admin")]), ), patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests - "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + "litellm.proxy.management_endpoints.team_endpoints.upsert_budget_and_membership", AsyncMock(), ), ): @@ -13917,7 +13919,7 @@ def _member_update_patches(team_snapshot: LiteLLM_TeamTable): ), ), patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests - "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + "litellm.proxy.management_endpoints.team_endpoints.upsert_budget_and_membership", AsyncMock(), ), ) @@ -14439,7 +14441,12 @@ def _wire_update_team(stack, existing_metadata): stack.enter_context(patch("litellm.proxy.proxy_server.user_api_key_cache")) stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj")) stack.enter_context(patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin")) - stack.enter_context(patch("litellm.proxy.management_endpoints.team_endpoints._cache_team_object")) + stack.enter_context( + patch( + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", + new_callable=AsyncMock, + ) + ) existing_team = MagicMock() existing_team.metadata = existing_metadata @@ -14843,7 +14850,10 @@ async def test_update_team_syncs_access_group_assigned_team_ids_in_both_directio patch("litellm.proxy.proxy_server.user_api_key_cache"), patch("litellm.proxy.proxy_server.proxy_logging_obj"), patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), - patch("litellm.proxy.management_endpoints.team_endpoints._refresh_cached_team"), + patch( + "litellm.proxy.management_endpoints.team_endpoints.refresh_cached_team", + new_callable=AsyncMock, + ), patch( "litellm.proxy.management_helpers.access_group_team_sync.invalidate_access_group_cache", new_callable=AsyncMock, @@ -15037,7 +15047,7 @@ async def test_invalidate_access_group_cache_deletes_the_cached_object(): patch("litellm.proxy.proxy_server.user_api_key_cache", cache), patch("litellm.proxy.proxy_server.proxy_logging_obj", logging_obj), patch( - "litellm.proxy.management_helpers.access_group_team_sync._delete_cache_access_object", + "litellm.proxy.management_helpers.access_group_team_sync.delete_cache_access_object", new_callable=AsyncMock, ) as delete_cached, ): @@ -15569,7 +15579,7 @@ async def test_team_member_update_invalidates_team_member_spend_state_when_budge AsyncMock(return_value=team_info_response), ), patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests - "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + "litellm.proxy.management_endpoints.team_endpoints.upsert_budget_and_membership", AsyncMock(), ), ): @@ -15622,7 +15632,7 @@ async def test_team_member_update_skips_invalidation_when_no_budget_fields_sent( AsyncMock(return_value=team_info_response), ), patch( # test-quality-ok: no live DB here; matches this file's established convention for endpoint-logic unit tests - "litellm.proxy.management_endpoints.team_endpoints._upsert_budget_and_membership", + "litellm.proxy.management_endpoints.team_endpoints.upsert_budget_and_membership", AsyncMock(), ), ): diff --git a/tests/unit/proxy/management_endpoints/test_team_model_alias_merge.py b/tests/unit/proxy/management_endpoints/test_team_model_alias_merge.py index 7cdf60f043e..22e03328fe6 100644 --- a/tests/unit/proxy/management_endpoints/test_team_model_alias_merge.py +++ b/tests/unit/proxy/management_endpoints/test_team_model_alias_merge.py @@ -29,9 +29,7 @@ class TestTeamModelAddAtomicAppend: from litellm.proxy.management_endpoints.team_endpoints import team_model_add mock_request = MagicMock() - mock_user = UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user" - ) + mock_user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user") existing_team = MagicMock() existing_team.model_dump.return_value = { @@ -49,19 +47,15 @@ class TestTeamModelAddAtomicAppend: with ( patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( - "litellm.proxy.management_endpoints.team_endpoints._cache_team_object", + "litellm.proxy.management_endpoints.team_endpoints.cache_team_object", new_callable=AsyncMock, ), patch("litellm.proxy.proxy_server.user_api_key_cache"), patch("litellm.proxy.proxy_server.proxy_logging_obj"), ): - mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) mock_prisma.db.execute_raw = AsyncMock(return_value=None) - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=updated_team - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=updated_team) await team_model_add( data=TeamModelAddRequest(team_id="team-1", models=["new-model"]), diff --git a/tests/unit/proxy/management_endpoints/test_ui_sso.py b/tests/unit/proxy/management_endpoints/test_ui_sso.py index e053c4c94ea..4e128440703 100644 --- a/tests/unit/proxy/management_endpoints/test_ui_sso.py +++ b/tests/unit/proxy/management_endpoints/test_ui_sso.py @@ -2270,7 +2270,7 @@ class TestUISSO_FunctionsExistence: assert SSOAuthenticationHandler is not None # Check that the new _get_cli_state method exists - assert hasattr(SSOAuthenticationHandler, "_get_cli_state") + assert hasattr(SSOAuthenticationHandler, "get_cli_state") assert callable(SSOAuthenticationHandler._get_cli_state) @@ -3028,7 +3028,7 @@ class TestCLIKeyRegenerationFlow: return_value="https://proxy.example.com/sso/callback", ), patch( - "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler._get_cli_state", + "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.get_cli_state", return_value=None, ) as mock_get_cli_state, ): @@ -8161,7 +8161,7 @@ class TestPKCEStateCookieBinding: ), patch.object( SSOAuthenticationHandler, - "_pkce_token_exchange", + "pkce_token_exchange", AsyncMock( return_value={ "access_token": "tok", @@ -8173,7 +8173,7 @@ class TestPKCEStateCookieBinding: ), patch.object( SSOAuthenticationHandler, - "_delete_pkce_verifier", + "delete_pkce_verifier", AsyncMock(), ), patch("fastapi_sso.sso.base.DiscoveryDocument"), @@ -8824,7 +8824,7 @@ async def test_pkce_arm_captures_sso_assertion(): ), patch.object( SSOAuthenticationHandler, - "_pkce_token_exchange", + "pkce_token_exchange", AsyncMock( return_value={ "access_token": "tok", @@ -8835,7 +8835,7 @@ async def test_pkce_arm_captures_sso_assertion(): } ), ), - patch.object(SSOAuthenticationHandler, "_delete_pkce_verifier", AsyncMock()), + patch.object(SSOAuthenticationHandler, "delete_pkce_verifier", AsyncMock()), patch("fastapi_sso.sso.base.DiscoveryDocument"), patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()), patch.dict( diff --git a/tests/unit/proxy/management_helpers/test_management_helpers_utils.py b/tests/unit/proxy/management_helpers/test_management_helpers_utils.py index 82eafc70077..b450a4dfd02 100644 --- a/tests/unit/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/unit/proxy/management_helpers/test_management_helpers_utils.py @@ -899,7 +899,7 @@ class _FakeDb: async def test_team_update_reaches_inherited_members_but_not_overridden_ones(): from litellm.proxy._types import LitellmUserRoles from litellm.proxy.auth.auth_checks import _check_team_member_budget - from litellm.proxy.management_endpoints.common_utils import _upsert_budget_and_membership + from litellm.proxy.management_endpoints.common_utils import upsert_budget_and_membership from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler from litellm.proxy.utils import ProxyLogging @@ -922,7 +922,7 @@ async def test_team_update_reaches_inherited_members_but_not_overridden_ones(): default_team_budget_id=default_budget.budget_id, ) - await _upsert_budget_and_membership( + await upsert_budget_and_membership( db, team_id=team_id, user_id="overridden", diff --git a/tests/unit/proxy/management_helpers/test_object_permission_utils.py b/tests/unit/proxy/management_helpers/test_object_permission_utils.py index 2fba39b30f6..33557713162 100644 --- a/tests/unit/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/unit/proxy/management_helpers/test_object_permission_utils.py @@ -19,7 +19,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( _extract_requested_mcp_access_groups, _extract_requested_mcp_server_ids, _resolve_team_allowed_mcp_servers, - _set_object_permission, + set_object_permission, enforce_all_proxy_mcp_servers_grant_is_admin_only, prepare_object_permission_upsert, validate_key_mcp_servers_against_team, @@ -62,7 +62,7 @@ async def test_set_object_permission(): } # Call the function - result = await _set_object_permission( + result = await set_object_permission( data_json=data_json, prisma_client=mock_prisma_client ) @@ -116,7 +116,7 @@ async def test_set_object_permission_persists_mcp_tool_search_enabled(): }, } - await _set_object_permission(data_json=data_json, prisma_client=mock_prisma_client) + await set_object_permission(data_json=data_json, prisma_client=mock_prisma_client) created_data = ( mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[ @@ -139,7 +139,7 @@ async def test_set_object_permission_persists_skills(): "object_permission": LiteLLM_ObjectPermissionBase(skills=["private-skill"]).model_dump(), } - await _set_object_permission(data_json=data_json, prisma_client=mock_prisma_client) + await set_object_permission(data_json=data_json, prisma_client=mock_prisma_client) created_data = ( mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[ @@ -273,11 +273,11 @@ def _make_mock_mcp_manager(*existing_ids: str, servers=None): @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -295,11 +295,11 @@ async def test_validate_no_object_permission(mock_access_groups, mock_allow_all) new=_make_mock_mcp_manager("server-1", "server-2"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -320,11 +320,11 @@ async def test_validate_key_servers_within_team_scope( new=_make_mock_mcp_manager("server-1", "server-outside"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -348,11 +348,11 @@ async def test_validate_key_servers_outside_team_scope_raises( new=_make_mock_mcp_manager("server-1", "global-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value={"global-server"}, ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -373,11 +373,11 @@ async def test_validate_allow_all_keys_servers_always_allowed( new=_make_mock_mcp_manager("global-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value={"global-server"}, ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -395,11 +395,11 @@ async def test_validate_no_team_only_allow_all_keys(mock_access_groups, mock_all new=_make_mock_mcp_manager("private-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value={"global-server"}, ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -422,11 +422,11 @@ async def test_validate_no_team_non_global_server_raises( new=_make_mock_mcp_manager("private-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -448,11 +448,11 @@ async def test_validate_no_team_proxy_admin_can_assign_private_server( new=_make_mock_mcp_manager("private-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -471,11 +471,11 @@ async def test_validate_no_team_non_admin_private_server_still_raises( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -497,11 +497,11 @@ async def test_validate_no_team_proxy_admin_can_assign_access_group( new=_make_mock_mcp_manager("server-1", "server-outside"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -526,11 +526,11 @@ async def test_validate_proxy_admin_still_bounded_by_team_scope( new=_make_mock_mcp_manager("some-server"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -679,11 +679,11 @@ async def test_team_unified_access_group_without_servers_preserves_direct_grants new=_make_mock_mcp_manager("server-outside"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -707,11 +707,11 @@ async def test_validate_tool_permissions_validated_against_team( new=_make_mock_mcp_manager(), # empty registry — all IDs are stale ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -739,11 +739,11 @@ async def test_validate_stale_mcp_server_ids_are_silently_dropped( new=_make_mock_mcp_manager(), # empty registry — all IDs are stale ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -769,11 +769,11 @@ async def test_validate_stale_ids_in_mcp_tool_permissions_silently_dropped( new=_make_mock_mcp_manager(), # empty registry — all IDs are stale ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -804,11 +804,11 @@ async def test_validate_stale_mcp_server_ids_are_removed_from_object_permission( ), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -839,11 +839,11 @@ async def test_validate_mcp_server_alias_outside_team_scope_raises( ), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -891,11 +891,11 @@ def test_alias_grant_expands_on_other_region_after_save(): new=_make_mock_mcp_manager(), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -925,11 +925,11 @@ async def test_validate_db_mcp_server_alias_outside_team_scope_raises_when_regis @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -946,11 +946,11 @@ async def test_validate_access_groups_within_team_scope( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -970,11 +970,11 @@ async def test_validate_access_groups_outside_team_scope_raises( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -993,11 +993,11 @@ async def test_validate_access_groups_no_team_raises( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=["server-from-group"], ) @@ -1018,7 +1018,7 @@ async def test_validate_team_access_groups_resolve_to_servers( @pytest.mark.asyncio @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1037,7 +1037,7 @@ async def test_resolve_team_allowed_mcp_servers_string_tool_permissions( @pytest.mark.asyncio @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1059,7 +1059,7 @@ async def test_resolve_team_allowed_mcp_servers_dict_tool_permissions( @pytest.mark.asyncio @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1100,11 +1100,11 @@ async def test_resolve_team_all_proxy_sentinel_resolves_dynamically(mock_access_ new=_make_mock_mcp_manager("srv-x", "srv-y", "srv-z"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1131,11 +1131,11 @@ async def test_validate_key_scoped_to_server_added_after_team_all_proxy( new=_make_mock_mcp_manager("srv-x", "srv-z"), ) @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1274,11 +1274,11 @@ async def test_validate_search_tools_raises_when_not_subset(): @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1297,11 +1297,11 @@ async def test_personal_non_admin_cannot_assign_mcp_toolsets( @pytest.mark.asyncio @patch( - "litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids", + "litellm.proxy.management_helpers.object_permission_utils.get_allow_all_keys_server_ids", return_value=set(), ) @patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups", + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_mcp_servers_from_access_groups", new_callable=AsyncMock, return_value=[], ) @@ -1550,7 +1550,7 @@ async def test_set_object_permission_rejects_shared_alias_or_name_tool_permissio data_json = {"object_permission": {"mcp_tool_permissions": {identifier: ["read_wiki_structure"]}}} with pytest.raises(HTTPException) as exc_info: - await _set_object_permission(data_json=data_json, prisma_client=mock_prisma) + await set_object_permission(data_json=data_json, prisma_client=mock_prisma) assert exc_info.value.status_code == 400 assert all(server_id in str(exc_info.value.detail) for server_id in colliding_ids) diff --git a/tests/unit/proxy/management_helpers/test_team_metadata_validation.py b/tests/unit/proxy/management_helpers/test_team_metadata_validation.py index 26bcba775a5..6defb5e45b2 100644 --- a/tests/unit/proxy/management_helpers/test_team_metadata_validation.py +++ b/tests/unit/proxy/management_helpers/test_team_metadata_validation.py @@ -405,7 +405,7 @@ async def _drive_update(kind, existing_metadata, payload): patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), patch( - "litellm.proxy.management_endpoints.team_endpoints._refresh_cached_team", + "litellm.proxy.management_endpoints.team_endpoints.refresh_cached_team", new=AsyncMock(), ), ): diff --git a/tests/unit/proxy/openai_files_endpoint/test_batch_guardrails.py b/tests/unit/proxy/openai_files_endpoint/test_batch_guardrails.py index 8c8dc5d799f..2387caca33b 100644 --- a/tests/unit/proxy/openai_files_endpoint/test_batch_guardrails.py +++ b/tests/unit/proxy/openai_files_endpoint/test_batch_guardrails.py @@ -940,10 +940,10 @@ async def test_a_real_non_guardrail_enforcement_hook_drops_its_record(monkeypatc pins that, because the other tests raise their own exceptions. """ import litellm - from litellm.proxy.hooks.prompt_injection_detection import _OPTIONAL_PromptInjectionDetection + from litellm.proxy.hooks.prompt_injection_detection import OPTIONAL_PromptInjectionDetection from litellm.proxy._types import LiteLLMPromptInjectionParams - hook = _OPTIONAL_PromptInjectionDetection( + hook = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams(heuristics_check=True) ) monkeypatch.setattr(litellm, "callbacks", [hook]) diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py index 238970d5a8e..07e6c3ac89d 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py @@ -165,7 +165,7 @@ class TestAnthropicLoggingHandlerModelFallback: return mock_handler @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) @patch.object( AnthropicPassthroughLoggingHandler, "_create_anthropic_response_logging_payload" @@ -2310,7 +2310,7 @@ class TestAnthropicUsageOnlyFallback: @patch("litellm.completion_cost") @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) def test_handler_falls_back_when_assembly_returns_none( self, mock_assemble, mock_cost @@ -2336,7 +2336,7 @@ class TestAnthropicUsageOnlyFallback: @patch("litellm.completion_cost") @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) def test_handler_falls_back_when_assembly_raises(self, mock_assemble, mock_cost): import litellm @@ -2368,7 +2368,7 @@ class TestAnthropicUsageOnlyFallback: assert result["kwargs"]["response_cost"] == 0.0021 @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) def test_handler_returns_none_when_no_usage_recoverable(self, mock_assemble): # assembly fails AND the chunks carry no usage event, so there is nothing @@ -2392,10 +2392,10 @@ class TestAnthropicUsageOnlyFallback: assert result["kwargs"] == {} @patch.object( - AnthropicPassthroughLoggingHandler, "_build_usage_only_response_from_chunks" + AnthropicPassthroughLoggingHandler, "build_usage_only_response_from_chunks" ) @patch.object( - AnthropicPassthroughLoggingHandler, "_build_complete_streaming_response" + AnthropicPassthroughLoggingHandler, "build_complete_streaming_response" ) def test_handler_does_not_crash_when_usage_only_fallback_raises( self, mock_assemble, mock_fallback @@ -2839,11 +2839,7 @@ def test_handle_logging_anthropic_collected_chunks(all_chunks): "all_chunks": all_chunks, } - result = ( - AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks( - **sent_args - ) - ) + result = AnthropicPassthroughLoggingHandler.handle_logging_anthropic_collected_chunks(**sent_args) assert isinstance(result["result"], ModelResponse) print("result=", json.dumps(result, indent=4, default=str)) @@ -2857,7 +2853,7 @@ def test_build_complete_streaming_response(all_chunks): litellm_logging_obj = Mock() - result = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( + result = AnthropicPassthroughLoggingHandler.build_complete_streaming_response( all_chunks=all_chunks, model="claude-sonnet-4-5-20250929", litellm_logging_obj=litellm_logging_obj, diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py index 52c664a65a5..c30f6d9f15a 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_gemini_passthrough_logging_handler.py @@ -284,7 +284,7 @@ class TestGeminiPassthroughLoggingHandler: mock_logging_obj = self._create_mock_logging_obj() # Mock the _handle_logging method to capture the call - handler._handle_logging = AsyncMock() + handler.handle_logging = AsyncMock() # Mock httpx response mock_response = self._create_mock_httpx_response() @@ -316,8 +316,8 @@ class TestGeminiPassthroughLoggingHandler: assert mock_logging_obj.model_call_details["custom_llm_provider"] == "gemini" # Verify that _handle_logging was called with the correct kwargs - handler._handle_logging.assert_called_once() - call_kwargs = handler._handle_logging.call_args[1] + handler.handle_logging.assert_called_once() + call_kwargs = handler.handle_logging.call_args[1] assert call_kwargs["response_cost"] == 0.000050 assert call_kwargs["model"] == "gemini-2.0-flash" assert call_kwargs["custom_llm_provider"] == "gemini" 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 1115b55e027..86a670bcb3e 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 @@ -1555,7 +1555,7 @@ class TestOpenAIPassthroughIntegration: ) # Mock the _handle_logging method to capture calls - self.handler._handle_logging = AsyncMock() + self.handler.handle_logging = AsyncMock() # Act result = await self.handler.pass_through_async_success_handler( @@ -1575,7 +1575,7 @@ class TestOpenAIPassthroughIntegration: ) # Assert - Should call the base handler, not our OpenAI handler - self.handler._handle_logging.assert_called_once() + self.handler.handle_logging.assert_called_once() @patch("litellm.cost_calculator.default_image_cost_calculator") def test_calculate_image_generation_cost(self, mock_image_cost_calculator): diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py index 7961d2a911b..7a411c95afa 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_typesafe_passthrough_logging_handler.py @@ -213,9 +213,9 @@ async def test_oss_gateway_accounts_for_checkpoint_usage_and_registered_cost( from litellm.caching.caching import DualCache from litellm.exceptions import BudgetExceededError - from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter + from litellm.proxy.hooks.model_max_budget_limiter import PROXY_VirtualKeyModelMaxBudgetLimiter - budget_limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) + budget_limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) assert await budget_limiter.is_key_within_model_budget(auth, f"{provider}/{requested}") await budget_limiter.async_log_success_event(logged, None, start, datetime.now()) with pytest.raises(BudgetExceededError): diff --git a/tests/unit/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py b/tests/unit/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py index 44533f35c72..404ed0f5800 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py +++ b/tests/unit/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py @@ -17,7 +17,7 @@ import litellm from litellm.caching.dual_cache import DualCache from litellm.proxy._lazy_features import LAZY_FEATURES from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import _cache_key_object +from litellm.proxy.auth.auth_checks import cache_key_object from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( _websocket_relay, deepgram_listen_websocket_route, @@ -365,7 +365,7 @@ def test_deepgram_listen_authenticates_the_litellm_key_and_relays_to_deepgram(mo async def _cache_restricted_key(virtual_key: str, models: list[str]) -> DualCache: cache = DualCache() - await _cache_key_object( + await cache_key_object( hashed_token=hash_token(virtual_key), user_api_key_obj=UserAPIKeyAuth(token=hash_token(virtual_key), models=models), user_api_key_cache=cache, 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 14fd60c9793..1887021a53a 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 @@ -1393,7 +1393,7 @@ class TestBedrockLLMProxyRoute: with ( patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.read_request_body", return_value=mock_request_body, ), patch( @@ -1435,7 +1435,7 @@ class TestBedrockLLMProxyRoute: with ( patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.read_request_body", return_value=mock_request_body, ), patch( @@ -2493,7 +2493,7 @@ class TestForwardHeaders: with ( patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.read_request_body", return_value=mock_request_body, ), patch( @@ -2591,7 +2591,7 @@ class TestForwardHeaders: with ( patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.read_request_body", return_value=mock_request_body, ), patch( @@ -2678,7 +2678,7 @@ class TestForwardHeaders: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials" ) as mock_get_creds, patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints._read_request_body", + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.read_request_body", return_value={"messages": [{"role": "user", "content": "test"}]}, ), patch( @@ -2782,7 +2782,7 @@ class TestMilvusProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ) as mock_is_allowed, patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body" + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.safe_set_request_parsed_body" ) as mock_safe_set, patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" @@ -2996,7 +2996,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.safe_set_request_parsed_body"), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, patch.object(litellm, "vector_store_registry") as mock_vector_registry, ): @@ -3046,7 +3046,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.safe_set_request_parsed_body"), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, patch.object(litellm, "vector_store_registry") as mock_vector_registry, ): @@ -3101,7 +3101,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.safe_set_request_parsed_body"), patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -4816,7 +4816,7 @@ class TestAzureProxyRouteCrossIndexAuthorization: new=AsyncMock(), ), patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler.base_openai_pass_through_handler", new=AsyncMock(return_value=Response()), ), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, @@ -4862,7 +4862,7 @@ class TestAzureProxyRouteCrossIndexAuthorization: return_value="azure-key", ), patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler.base_openai_pass_through_handler", new=AsyncMock(return_value=Response()), ) as mock_handler, patch.object(litellm, "vector_store_index_registry") as mock_index_registry, @@ -4922,7 +4922,7 @@ class TestAzureProxyRouteServiceLevelIndexCreate: return_value="https://svc.search.windows.net", ), patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler.base_openai_pass_through_handler", new=AsyncMock(return_value=Response()), ) as mock_handler, ): @@ -4954,7 +4954,7 @@ class TestAzureProxyRouteServiceLevelIndexCreate: return_value="azure-key", ), patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler.base_openai_pass_through_handler", new=AsyncMock(return_value=Response()), ) as mock_handler, ): @@ -7542,16 +7542,16 @@ class TestTypeSafePassthroughRoute: ) -> None: from litellm.caching.caching import DualCache from litellm.proxy import proxy_server - from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck + from litellm.proxy.hooks.cache_control_check import PROXY_CacheControlCheck from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, get_request_stash, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging cache: Final = DualCache() - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) - monkeypatch.setattr(litellm, "callbacks", list((limiter, _PROXY_CacheControlCheck()))) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) + monkeypatch.setattr(litellm, "callbacks", list((limiter, PROXY_CacheControlCheck()))) monkeypatch.setattr(proxy_server, "proxy_logging_obj", ProxyLogging(user_api_key_cache=cache)) monkeypatch.setenv("OPENROUTER_API_KEY", "openrouter-test-key") monkeypatch.setenv("OPENROUTER_API_BASE", "https://typesafe.example/base") @@ -7719,12 +7719,12 @@ class TestOssDecisionPassthroughRoute: self, client: TestClient, monkeypatch: pytest.MonkeyPatch, metadata_slot: str, provider: str, checkpoint: str ) -> None: from litellm.integrations.custom_logger import CustomLogger - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import InternalUsageCache from litellm.proxy.proxy_server import app cache: Final = DualCache() - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) auth: Final = UserAPIKeyAuth( api_key="oss-native-rpm", metadata={"model_rpm_limit": {f"{provider}/{checkpoint}": 1}}, ) diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 4142e5b94c6..cd279a6ed4b 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -580,7 +580,7 @@ async def test_custom_passthrough_predict_path_logs_via_generic_handler(): ) handler = PassThroughEndpointLogging() - handler._handle_logging = AsyncMock() + handler.handle_logging = AsyncMock() mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj) mock_logging_obj.model_call_details = {} @@ -612,8 +612,8 @@ async def test_custom_passthrough_predict_path_logs_via_generic_handler(): ) mock_vertex_handler.assert_not_called() - handler._handle_logging.assert_awaited_once() - logged_object = handler._handle_logging.call_args.kwargs["standard_logging_response_object"] + handler.handle_logging.assert_awaited_once() + logged_object = handler.handle_logging.call_args.kwargs["standard_logging_response_object"] assert logged_object == {"response": '{"forecast": [1, 2, 3]}'} @@ -1073,7 +1073,7 @@ async def test_pass_through_success_handler_with_cost_per_request(): mock_logging_obj.model_call_details = {} # Mock the _handle_logging method to capture the call - handler._handle_logging = AsyncMock() + handler.handle_logging = AsyncMock() # Mock httpx response mock_response = MagicMock(spec=httpx.Response) @@ -1108,8 +1108,8 @@ async def test_pass_through_success_handler_with_cost_per_request(): assert mock_logging_obj.model_call_details["response_cost"] == 1.25 # Verify that _handle_logging was called with the correct kwargs - handler._handle_logging.assert_called_once() - call_kwargs = handler._handle_logging.call_args[1] + handler.handle_logging.assert_called_once() + call_kwargs = handler.handle_logging.call_args[1] assert call_kwargs["response_cost"] == 1.25 @@ -3913,11 +3913,11 @@ async def _drive_pass_through_block(raised_exception): patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), patch(f"{_PT_MODULE}.verbose_proxy_logger", logger), patch( - f"{_PT_MODULE}._read_request_body", + f"{_PT_MODULE}.read_request_body", new_callable=AsyncMock, return_value={}, ), - patch(f"{_PT_MODULE}._safe_get_request_headers", return_value={}), + patch(f"{_PT_MODULE}.safe_get_request_headers", return_value={}), patch( "litellm.proxy.pass_through_endpoints.passthrough_guardrails." "PassthroughGuardrailHandler.collect_guardrails", @@ -6708,7 +6708,7 @@ def _passthrough_kwargs_for_reservation( async def _track_cost_for_passthrough_kwargs(kwargs: dict) -> AsyncMock: from datetime import datetime - from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger callback_kwargs = { **kwargs, @@ -6731,7 +6731,7 @@ async def _track_cost_for_passthrough_kwargs(kwargs: dict) -> AsyncMock: mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock() - await _ProxyDBLogger()._PROXY_track_cost_callback( + await ProxyDBLogger()._PROXY_track_cost_callback( kwargs=callback_kwargs, completion_response=None, start_time=datetime.now(), @@ -7033,10 +7033,10 @@ async def test_user_defined_passthrough_is_neither_tracked_nor_enforced(metadata from litellm.caching.caching import DualCache from litellm.proxy.auth.auth_utils import get_model_from_request - from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter + from litellm.proxy.hooks.model_max_budget_limiter import PROXY_VirtualKeyModelMaxBudgetLimiter budget: Final = {"managed-model": {"budget_limit": 0.1, "time_period": "1d"}} - limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) + limiter: Final = PROXY_VirtualKeyModelMaxBudgetLimiter(DualCache()) auth: Final = UserAPIKeyAuth( api_key="custom-key", token="custom-key", team_id="shared-team", team_model_max_budget=budget, ) diff --git a/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py b/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py index 09987b2781c..e0687636263 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py +++ b/tests/unit/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py @@ -147,8 +147,8 @@ async def _drive(response_text: str): patch("litellm.proxy.proxy_server.llm_router", None), patch(f"{_PT_MOD}.pass_through_endpoint_logging", mock_pt_logging), patch(f"{_PT_MOD}.get_async_httpx_client", return_value=mock_async_client_obj), - patch(f"{_PT_MOD}._read_request_body", new_callable=AsyncMock, return_value={}), - patch(f"{_PT_MOD}._safe_get_request_headers", return_value={}), + patch(f"{_PT_MOD}.read_request_body", new_callable=AsyncMock, return_value={}), + patch(f"{_PT_MOD}.safe_get_request_headers", return_value={}), patch(_COLLECT, return_value=["block-demo"]), ] try: diff --git a/tests/unit/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py b/tests/unit/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py index c7696079adc..7b06905c8aa 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py +++ b/tests/unit/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py @@ -89,8 +89,8 @@ def _common_patches(mock_proxy_logging, mock_response): patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), patch(f"{_PT_MOD}.pass_through_endpoint_logging", mock_pt_logging), patch(f"{_PT_MOD}.get_async_httpx_client", return_value=mock_async_client_obj), - patch(f"{_PT_MOD}._read_request_body", new_callable=AsyncMock, return_value={}), - patch(f"{_PT_MOD}._safe_get_request_headers", return_value={}), + patch(f"{_PT_MOD}.read_request_body", new_callable=AsyncMock, return_value={}), + patch(f"{_PT_MOD}.safe_get_request_headers", return_value={}), ] stack = ExitStack() 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 e6b19f4eec1..b278d7dea01 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 @@ -55,7 +55,7 @@ async def test_chunk_processor_logs_on_normal_completion(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: received = [] @@ -88,7 +88,7 @@ async def test_chunk_processor_logs_on_client_disconnect(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: gen = PassThroughStreamingHandler.chunk_processor( @@ -126,7 +126,7 @@ async def test_chunk_processor_does_not_schedule_success_logging_for_upstream_er with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: received = [] @@ -156,7 +156,7 @@ async def test_chunk_processor_does_not_schedule_logging_when_no_chunks(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ) as mock_route: received = [] @@ -193,7 +193,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker(): with ( patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ), patch.object( @@ -235,7 +235,7 @@ async def test_chunk_processor_routes_logging_through_logging_worker_on_disconne with ( patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ), patch.object( @@ -287,7 +287,7 @@ async def test_chunk_processor_stamps_completion_start_time_on_first_chunk(): with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ): received = [] @@ -326,7 +326,7 @@ async def test_chunk_processor_does_not_reset_completion_start_time_on_later_chu with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ): async for _ in PassThroughStreamingHandler.chunk_processor( @@ -361,7 +361,7 @@ async def test_chunk_processor_stamps_completion_start_time_on_cost_injection_pa try: with patch.object( PassThroughStreamingHandler, - "_route_streaming_logging_to_handler", + "route_streaming_logging_to_handler", new=AsyncMock(), ): async for _ in PassThroughStreamingHandler.chunk_processor( diff --git a/tests/unit/proxy/proxy_server/test_background_health.py b/tests/unit/proxy/proxy_server/test_background_health.py index b15349d6705..bb2fac403a7 100644 --- a/tests/unit/proxy/proxy_server/test_background_health.py +++ b/tests/unit/proxy/proxy_server/test_background_health.py @@ -178,7 +178,7 @@ async def test_schedule_background_health_check_db_save_creates_task(monkeypatch import litellm.proxy.health_endpoints._health_endpoints as he - monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + monkeypatch.setattr(he, "save_background_health_checks_to_db", _fake_save) prisma_client = MagicMock() shared_manager = SimpleNamespace(pod_id="pod-xyz") @@ -227,7 +227,7 @@ async def test_schedule_background_health_check_db_save_invalid_no_event_loop_ra import litellm.proxy.health_endpoints._health_endpoints as he - monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + monkeypatch.setattr(he, "save_background_health_checks_to_db", _fake_save) def _broken_create_task(_coro): raise RuntimeError("no running event loop") @@ -261,7 +261,7 @@ def _capture_saves(monkeypatch, persisted=True): import litellm.proxy.health_endpoints._health_endpoints as he - monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + monkeypatch.setattr(he, "save_background_health_checks_to_db", _fake_save) return saves @@ -271,7 +271,7 @@ def _cancel_during_save(monkeypatch): import litellm.proxy.health_endpoints._health_endpoints as he - monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + monkeypatch.setattr(he, "save_background_health_checks_to_db", _fake_save) def _schedule_with(lock_manager): diff --git a/tests/unit/proxy/proxy_server/test_lifecycle.py b/tests/unit/proxy/proxy_server/test_lifecycle.py index ba5501315d9..69f9c72c9c8 100644 --- a/tests/unit/proxy/proxy_server/test_lifecycle.py +++ b/tests/unit/proxy/proxy_server/test_lifecycle.py @@ -33,7 +33,7 @@ from typing_extensions import TypedDict import litellm.proxy.proxy_server as ps from litellm.proxy.proxy_server import ( ProxyStartupEvent, - _initialize_shared_aiohttp_session, + initialize_shared_aiohttp_session, _resolve_pydantic_type, _resolve_typed_dict_type, cleanup_router_config_variables, @@ -349,7 +349,7 @@ async def test_flush_spend_counters_on_shutdown_logs_and_swallows_commit_errors( async def test_initialize_shared_aiohttp_session_returns_client_session(): from aiohttp import ClientSession - session = await _initialize_shared_aiohttp_session() + session = await initialize_shared_aiohttp_session() try: observed = { "is_client_session": isinstance(session, ClientSession), @@ -381,7 +381,7 @@ async def test_initialize_shared_aiohttp_session_aiohttp_missing_returns_none_on return real_import(name, *args, **kwargs) monkeypatch.setattr(builtins, "__import__", _raise_for_aiohttp) - result = await _initialize_shared_aiohttp_session() + result = await initialize_shared_aiohttp_session() assert result is None @@ -914,7 +914,7 @@ def test_otel_global_provider_published_after_callback_init(): """ wrapped = getattr(proxy_startup_event, "__wrapped__", proxy_startup_event) source = inspect.getsource(wrapped) - init_pos = source.find("_initialize_startup_logging(") + init_pos = source.find("ProxyStartupEvent.initialize_startup_logging(") publish_pos = source.find("publish_global_otel_v2_provider(") assert init_pos != -1, "callback init call not found in proxy_startup_event" assert publish_pos != -1, "OTEL global publish not found in proxy_startup_event" @@ -927,7 +927,7 @@ def test_otel_global_provider_published_after_callback_init(): def test_startup_warns_for_global_budget_without_database(caplog): with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_budget_without_db(max_budget=100.0, prisma_client=None) + ProxyStartupEvent.warn_budget_without_db(max_budget=100.0, prisma_client=None) assert "litellm.max_budget=100.0" in caplog.text assert "will NOT be enforced" in caplog.text @@ -936,7 +936,7 @@ def test_startup_warns_for_global_budget_without_database(caplog): def test_startup_does_not_warn_for_global_budget_with_database(caplog): with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_budget_without_db(max_budget=100.0, prisma_client=MagicMock()) + ProxyStartupEvent.warn_budget_without_db(max_budget=100.0, prisma_client=MagicMock()) assert "litellm.max_budget" not in caplog.text @@ -944,7 +944,7 @@ def test_startup_does_not_warn_for_global_budget_with_database(caplog): @pytest.mark.parametrize("max_budget", [0, None]) def test_startup_does_not_warn_without_global_budget(caplog, max_budget): with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_budget_without_db(max_budget=max_budget, prisma_client=None) + ProxyStartupEvent.warn_budget_without_db(max_budget=max_budget, prisma_client=None) assert "litellm.max_budget" not in caplog.text @@ -1001,7 +1001,7 @@ def test_proxy_startup_event_warns_for_global_budget_without_database(): wrapped = getattr(proxy_startup_event, "__wrapped__", proxy_startup_event) source = inspect.getsource(wrapped) budget_check_pos = source.find("if prisma_client is not None and litellm.max_budget > 0:") - warn_pos = source.find("_warn_budget_without_db(") + warn_pos = source.find("warn_budget_without_db(") next_startup_section_pos = source.find( "await ProxyStartupEvent.initialize_scheduled_background_jobs(", budget_check_pos, @@ -1089,7 +1089,7 @@ async def test_scorer_baseline_upgrade_preserves_existing_routers_and_is_not_ref @pytest.mark.asyncio async def test_tuning_baseline_waits_for_a_complete_db_model_census(monkeypatch): prisma_client = MagicMock() - monkeypatch.setattr(ps.proxy_config, "_get_models_from_db", AsyncMock(return_value=None)) + monkeypatch.setattr(ps.proxy_config, "get_models_from_db", AsyncMock(return_value=None)) result = await ProxyStartupEvent.enforce_heuristic_v1_tuning_baseline( prisma_client=prisma_client, diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index 903557194ec..d1b80627b3a 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -3883,7 +3883,7 @@ async def test_ProxyConfig_add_deployment_applies_db_router_settings(monkeypatch return {} monkeypatch.setattr(pc, "get_config", fake_get_config) - monkeypatch.setattr(pc, "_get_models_from_db", AsyncMock(return_value=[])) + monkeypatch.setattr(pc, "get_models_from_db", AsyncMock(return_value=[])) monkeypatch.setattr(pc, "_init_non_llm_objects_in_db", AsyncMock()) monkeypatch.setattr(proxy_server, "prefetch_config_params", AsyncMock()) monkeypatch.setattr(proxy_server, "get_config_param", AsyncMock(return_value=None)) @@ -3963,7 +3963,7 @@ async def test_ProxyConfig_add_deployment_loads_db_credentials_before_reconcilin async def install_models(new_models: object, proxy_logging_obj: object) -> None: installed(credential=CredentialAccessor.get_credential_values("openai-cred")) - monkeypatch.setattr(pc, "_get_models_from_db", read_models_while_a_credential_lands) + monkeypatch.setattr(pc, "get_models_from_db", read_models_while_a_credential_lands) monkeypatch.setattr(pc, "_update_llm_router", install_models) await pc.add_deployment(prisma_client=fake_prisma, proxy_logging_obj=MagicMock()) @@ -3987,7 +3987,7 @@ async def test_ProxyConfig_add_deployment_loads_db_credentials_even_when_models_ _stub_add_deployment_collaborators(monkeypatch, pc, fake_prisma) monkeypatch.setattr(proxy_server, "general_settings", {"supported_db_objects": ["mcp"]}) models_fetch = AsyncMock(return_value=[]) - monkeypatch.setattr(pc, "_get_models_from_db", models_fetch) + monkeypatch.setattr(pc, "get_models_from_db", models_fetch) await pc.add_deployment(prisma_client=fake_prisma, proxy_logging_obj=MagicMock()) diff --git a/tests/unit/proxy/proxy_server/test_routes_chat_completions.py b/tests/unit/proxy/proxy_server/test_routes_chat_completions.py index b186bb5ef5e..e124bde8120 100644 --- a/tests/unit/proxy/proxy_server/test_routes_chat_completions.py +++ b/tests/unit/proxy/proxy_server/test_routes_chat_completions.py @@ -37,9 +37,7 @@ HAPPY_RESPONSE = { def patched_chat(monkeypatch): """Stub chat-completions pipeline at ProxyBaseLLMRequestProcessing.""" monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) - monkeypatch.setattr( - proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) - ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())) async def _fake_process(self, *args, **kwargs): return dict(HAPPY_RESPONSE) @@ -56,9 +54,7 @@ def patched_chat(monkeypatch): def patched_chat_error(monkeypatch): """Variant that makes the pipeline raise -> 400 via _handle_llm_api_exception.""" monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) - monkeypatch.setattr( - proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) - ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())) from litellm.proxy._types import ProxyException @@ -66,9 +62,7 @@ def patched_chat_error(monkeypatch): raise ValueError("boom") async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj): - return ProxyException( - message="boom", type="bad_request_error", param="model", code=400 - ) + return ProxyException(message="boom", type="bad_request_error", param="model", code=400) monkeypatch.setattr( common_request_processing.ProxyBaseLLMRequestProcessing, @@ -77,7 +71,7 @@ def patched_chat_error(monkeypatch): ) monkeypatch.setattr( common_request_processing.ProxyBaseLLMRequestProcessing, - "_handle_llm_api_exception", + "handle_llm_api_exception", _handler, ) yield diff --git a/tests/unit/proxy/proxy_server/test_routes_config.py b/tests/unit/proxy/proxy_server/test_routes_config.py index c470fa96d42..4c8265b82c3 100644 --- a/tests/unit/proxy/proxy_server/test_routes_config.py +++ b/tests/unit/proxy/proxy_server/test_routes_config.py @@ -1518,7 +1518,7 @@ def test_get_config_callbacks_excludes_internal_runtime_callbacks(client, auth_a from litellm.integrations.s3_v2 import S3Logger from litellm.integrations.sqs import SQSLogger from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import VectorStorePreCallHook - from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck + from litellm.proxy.hooks.cache_control_check import PROXY_CacheControlCheck from litellm.router import Router class _InventoryTestGuardrail(CustomGuardrail): @@ -1546,7 +1546,7 @@ def test_get_config_callbacks_excludes_internal_runtime_callbacks(client, auth_a litellm, "callbacks", [ - _PROXY_CacheControlCheck(), + PROXY_CacheControlCheck(), PROXY_LiteLLMManagedFiles(internal_usage_cache=MagicMock(), prisma_client=MagicMock()), ServiceLogging(), VectorStorePreCallHook(), diff --git a/tests/unit/proxy/proxy_server/test_routes_embeddings.py b/tests/unit/proxy/proxy_server/test_routes_embeddings.py index 98249cb5ad5..f511938ee23 100644 --- a/tests/unit/proxy/proxy_server/test_routes_embeddings.py +++ b/tests/unit/proxy/proxy_server/test_routes_embeddings.py @@ -31,9 +31,7 @@ def patched_embedding(monkeypatch): router.model_names = ["text-embedding-ada-002"] router.get_deployment_by_model_group_name = MagicMock(return_value=None) monkeypatch.setattr(proxy_server, "llm_router", router) - monkeypatch.setattr( - proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) - ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())) async def _fake_process(self, *args, **kwargs): return dict(HAPPY_RESPONSE) @@ -51,9 +49,7 @@ def embedding_pipeline_raises(monkeypatch): router = MagicMock() router.model_names = [] monkeypatch.setattr(proxy_server, "llm_router", router) - monkeypatch.setattr( - proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) - ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())) from litellm.proxy._types import ProxyException @@ -61,9 +57,7 @@ def embedding_pipeline_raises(monkeypatch): raise ValueError("boom") async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj, version=None): - return ProxyException( - message="boom", type="bad_request_error", param="model", code=400 - ) + return ProxyException(message="boom", type="bad_request_error", param="model", code=400) monkeypatch.setattr( common_request_processing.ProxyBaseLLMRequestProcessing, @@ -72,7 +66,7 @@ def embedding_pipeline_raises(monkeypatch): ) monkeypatch.setattr( common_request_processing.ProxyBaseLLMRequestProcessing, - "_handle_llm_api_exception", + "handle_llm_api_exception", _handler, ) yield diff --git a/tests/unit/proxy/proxy_server/test_routes_invitation.py b/tests/unit/proxy/proxy_server/test_routes_invitation.py index 5b54a63d8a2..35f2991996a 100644 --- a/tests/unit/proxy/proxy_server/test_routes_invitation.py +++ b/tests/unit/proxy/proxy_server/test_routes_invitation.py @@ -101,8 +101,8 @@ def test_invitation_new_non_admin_forbidden(client, auth_as, monkeypatch, mock_p return False # Patch at the proxy_server import site (used by the route). - monkeypatch.setattr(ps, "_user_has_admin_privileges", _no_privileges) - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _no_privileges) + monkeypatch.setattr(ps, "user_has_admin_privileges", _no_privileges) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _no_privileges) with auth_as(LitellmUserRoles.INTERNAL_USER): response = client.post("/invitation/new", json={"user_id": "user-target"}) @@ -186,7 +186,7 @@ def test_invitation_info_not_admin_forbidden(client, auth_as, monkeypatch, mock_ monkeypatch.setattr(ps, "prisma_client", mock_prisma) # _user_has_admin_view is referenced from proxy_server's import. - monkeypatch.setattr(ps, "_user_has_admin_view", lambda u: False) + monkeypatch.setattr(ps, "user_api_key_has_admin_view", lambda u: False) with auth_as(LitellmUserRoles.INTERNAL_USER): response = client.get("/invitation/info", params={"invitation_id": "inv-xyz"}) @@ -339,7 +339,7 @@ def test_invitation_delete_non_admin_forbidden( async def _no_privileges(**kwargs): return False - monkeypatch.setattr(ps, "_user_has_admin_privileges", _no_privileges) + monkeypatch.setattr(ps, "user_has_admin_privileges", _no_privileges) with auth_as(LitellmUserRoles.INTERNAL_USER): response = client.post( diff --git a/tests/unit/proxy/proxy_server/test_routes_model_info.py b/tests/unit/proxy/proxy_server/test_routes_model_info.py index 5175d92084c..fae4a6470cf 100644 --- a/tests/unit/proxy/proxy_server/test_routes_model_info.py +++ b/tests/unit/proxy/proxy_server/test_routes_model_info.py @@ -648,7 +648,7 @@ def model_group_info_router(monkeypatch): monkeypatch.setattr(proxy_server, "general_settings", {}) monkeypatch.setattr(proxy_server, "prisma_client", None) monkeypatch.setattr(proxy_server, "user_api_key_cache", None) - monkeypatch.setattr(proxy_server, "_get_model_group_info", model_group_info) + monkeypatch.setattr(proxy_server, "get_model_group_info", model_group_info) from litellm.proxy.agent_endpoints import model_list_helpers diff --git a/tests/unit/proxy/proxy_server/test_streaming_helpers.py b/tests/unit/proxy/proxy_server/test_streaming_helpers.py index 69fa195e9d6..09d71cb4ead 100644 --- a/tests/unit/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/unit/proxy/proxy_server/test_streaming_helpers.py @@ -678,7 +678,7 @@ def _patch_logging_flags(monkeypatch, needs_wrap=False, needs_per_chunk=False): # touching real logging globals. monkeypatch.setattr( ps.ProxyLogging, - "_fire_deferred_stream_logging", + "fire_deferred_stream_logging", staticmethod(lambda request_data: None), ) diff --git a/tests/unit/proxy/public_endpoints/test_public_endpoints.py b/tests/unit/proxy/public_endpoints/test_public_endpoints.py index 99a76d59448..30538c1167f 100644 --- a/tests/unit/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/unit/proxy/public_endpoints/test_public_endpoints.py @@ -548,11 +548,11 @@ def test_public_model_hub_with_healthy_model(): with ( patch("litellm.public_model_groups", ["gpt-3.5-turbo"]), - patch("litellm.proxy.proxy_server._get_model_group_info") as mock_get_info, + patch("litellm.proxy.proxy_server.get_model_group_info") as mock_get_info, patch("litellm.proxy.proxy_server.llm_router", mock_llm_router), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( - "litellm.proxy.health_endpoints._health_endpoints._convert_health_check_to_dict" + "litellm.proxy.health_endpoints._health_endpoints.convert_health_check_to_dict" ) as mock_convert, ): @@ -606,11 +606,11 @@ def test_public_model_hub_with_unhealthy_model(): with ( patch("litellm.public_model_groups", ["gpt-4"]), - patch("litellm.proxy.proxy_server._get_model_group_info") as mock_get_info, + patch("litellm.proxy.proxy_server.get_model_group_info") as mock_get_info, patch("litellm.proxy.proxy_server.llm_router", mock_llm_router), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( - "litellm.proxy.health_endpoints._health_endpoints._convert_health_check_to_dict" + "litellm.proxy.health_endpoints._health_endpoints.convert_health_check_to_dict" ) as mock_convert, ): @@ -655,7 +655,7 @@ def test_public_model_hub_without_health_check(): with ( patch("litellm.public_model_groups", ["claude-3"]), - patch("litellm.proxy.proxy_server._get_model_group_info") as mock_get_info, + patch("litellm.proxy.proxy_server.get_model_group_info") as mock_get_info, patch("litellm.proxy.proxy_server.llm_router", mock_llm_router), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), ): @@ -737,11 +737,11 @@ def test_public_model_hub_mixed_health_statuses(): with ( patch("litellm.public_model_groups", ["gpt-3.5-turbo", "gpt-4", "claude-3"]), - patch("litellm.proxy.proxy_server._get_model_group_info") as mock_get_info, + patch("litellm.proxy.proxy_server.get_model_group_info") as mock_get_info, patch("litellm.proxy.proxy_server.llm_router", mock_llm_router), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( - "litellm.proxy.health_endpoints._health_endpoints._convert_health_check_to_dict" + "litellm.proxy.health_endpoints._health_endpoints.convert_health_check_to_dict" ) as mock_convert, ): diff --git a/tests/unit/proxy/response_api_endpoints/test_endpoints.py b/tests/unit/proxy/response_api_endpoints/test_endpoints.py index 456c3609419..a8aff3ffff0 100644 --- a/tests/unit/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/unit/proxy/response_api_endpoints/test_endpoints.py @@ -137,7 +137,7 @@ async def test_responses_api_background_polling_rejects_missing_input(): async def return_exception(*, e: Exception, **kwargs: object) -> Exception: return e - processor._handle_llm_api_exception = AsyncMock(side_effect=return_exception) + processor.handle_llm_api_exception = AsyncMock(side_effect=return_exception) processor.common_processing_pre_call_logic = AsyncMock(return_value=({"model": "gpt-4o"}, MagicMock())) async def receive(): @@ -960,11 +960,11 @@ class TestResponsesWSFirstFrameModelAuth: with ( patch( - "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", + "litellm.proxy.auth.user_api_key_auth.enforce_key_and_fallback_model_access", new_callable=AsyncMock, ) as mock_key_check, patch( - "litellm.proxy.auth.user_api_key_auth._run_centralized_common_checks", + "litellm.proxy.auth.user_api_key_auth.run_centralized_common_checks", new_callable=AsyncMock, ) as mock_common_checks, patch( @@ -1260,7 +1260,7 @@ def test_cursor_chat_completions_input_body_uses_responses_pipeline_and_strips_s from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body as real_read_request_body, + read_request_body as real_read_request_body, ) from litellm.types.llms.openai import ResponsesAPIResponse @@ -1297,7 +1297,7 @@ def test_cursor_chat_completions_input_body_uses_responses_pipeline_and_strips_s with ( patch.object(ps, "llm_router", mock_router), patch( - "litellm.proxy.response_api_endpoints.endpoints._read_request_body", + "litellm.proxy.response_api_endpoints.endpoints.read_request_body", side_effect=capturing_read_request_body, ), ): @@ -1407,9 +1407,9 @@ class TestCursorMessagesArmToolNormalization: seen = {} async def fake_chat_completion(request, fastapi_response, model, user_api_key_dict): - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body - seen["body"] = await _read_request_body(request=request) + seen["body"] = await read_request_body(request=request) return {"id": "chatcmpl-fake", "object": "chat.completion", "choices": []} app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(api_key=MASTER_KEY) @@ -1470,9 +1470,9 @@ class TestCursorMessagesArmToolNormalization: seen = {} async def fake_chat_completion(request, fastapi_response, model, user_api_key_dict): - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body - seen["body"] = await _read_request_body(request=request) + seen["body"] = await read_request_body(request=request) return {"id": "chatcmpl-fake", "object": "chat.completion", "choices": []} body = { @@ -1949,9 +1949,9 @@ class TestCursorModelSuffixResolutionEndToEnd: seen = {} async def fake_chat_completion(request, fastapi_response, model, user_api_key_dict): - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body - seen["body"] = await _read_request_body(request=request) + seen["body"] = await read_request_body(request=request) return {"id": "chatcmpl-fake", "object": "chat.completion", "choices": []} app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(api_key=MASTER_KEY) @@ -2031,7 +2031,7 @@ def _cursor_budget_auth_env(base_model: str, spend: float): from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.model_max_budget_limiter import ( VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX, - _PROXY_VirtualKeyModelMaxBudgetLimiter, + PROXY_VirtualKeyModelMaxBudgetLimiter, ) valid_token = UserAPIKeyAuth( @@ -2039,7 +2039,7 @@ def _cursor_budget_auth_env(base_model: str, spend: float): token="hashed-cursor-budget-token", model_max_budget={base_model: {"budget_limit": 0.00001, "time_period": "1d"}}, ) - limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter = PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) limiter.dual_cache.in_memory_cache.set_cache( key=f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{valid_token.token}:{base_model}:1d", value=spend, @@ -2127,14 +2127,14 @@ class TestCursorVariantResolvedBeforeAuth: def _run_with_recording_auth(self, mock_router, request_model: str): from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body from fastapi import Request bodies_seen_by_auth = [] async def recording_auth(request: Request) -> UserAPIKeyAuth: - bodies_seen_by_auth.append(await _read_request_body(request=request)) + bodies_seen_by_auth.append(await read_request_body(request=request)) return UserAPIKeyAuth(api_key="sk-test-cursor") async def fake_chat_completion(request, fastapi_response, model, user_api_key_dict): diff --git a/tests/unit/proxy/spend_tracking/test_search_api_logging.py b/tests/unit/proxy/spend_tracking/test_search_api_logging.py index 92108f19c8e..bcb079cf37b 100644 --- a/tests/unit/proxy/spend_tracking/test_search_api_logging.py +++ b/tests/unit/proxy/spend_tracking/test_search_api_logging.py @@ -18,7 +18,7 @@ import litellm from litellm import Router from litellm.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger +from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger from litellm.proxy.spend_tracking.spend_management_endpoints import view_spend_logs from litellm.proxy.utils import ProxyLogging, hash_token, update_spend from litellm.llms.base_llm.search.transformation import SearchResponse, SearchResult @@ -129,7 +129,7 @@ async def test_search_api_logging_and_cost_tracking(prisma_client): setattr(litellm.proxy.proxy_server, "proxy_logging_obj", proxy_logging_obj) # Call the track_cost_callback directly to simulate what happens after a search - proxy_db_logger = _ProxyDBLogger() + proxy_db_logger = ProxyDBLogger() # Simulate the kwargs that would be passed from the search endpoint request_id = "search_test_123" 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 2e3b3c13bb3..8cc8dccd13e 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py @@ -273,7 +273,7 @@ from litellm.proxy._types import ( SpendLogsPayload, UserAPIKeyAuth, ) -from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger +from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger from litellm.proxy.management.teams import authz as team_access from litellm.proxy.proxy_server import app from litellm.proxy.spend_tracking import spend_management_endpoints @@ -3684,7 +3684,7 @@ class TestSpendLogsPayload: @pytest.mark.asyncio async def test_spend_logs_payload_e2e(self): - litellm.callbacks = [_ProxyDBLogger(message_logging=False)] + litellm.callbacks = [ProxyDBLogger(message_logging=False)] # litellm.turn_on_debug() with ( diff --git a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py index d4b397e3fa2..0df6eecd4b9 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py @@ -27,10 +27,10 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import ( _get_session_id_for_spend_log, _get_spend_logs_metadata, _get_vector_store_request_for_spend_logs_payload, - _is_master_key, + is_master_key, _redact_logged_api_key, _redact_prompt_leaks_in_error_string, - _sanitize_error_information_for_spend_logs, + sanitize_error_information_for_spend_logs, _sanitize_guardrail_information_for_spend_logs, _sanitize_request_body_for_spend_logs_payload, _scrub_raw_model_from_error_information, @@ -1476,7 +1476,7 @@ def test_get_logging_payload_persists_no_raw_model_for_a_prompt_shaped_moderatio model=_RAW_MODEL_WITH_PROMPT, llm_provider="openai", ) - error_information: Final = _sanitize_error_information_for_spend_logs( + error_information: Final = sanitize_error_information_for_spend_logs( StandardLoggingPayloadSetup.get_error_information( original_exception=provider_rejection, traceback_str=( @@ -1643,7 +1643,7 @@ async def test_api_key_preserved_through_failure_hook_to_database(): If this test fails in CI/CD, the build MUST fail. """ from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger from litellm.proxy.utils import hash_token # Setup @@ -1727,7 +1727,7 @@ async def test_api_key_preserved_through_failure_hook_to_database(): exception = Exception("BadRequestError: Invalid parameter 'invalid_param'") # Execute the ACTUAL failure hook code path - logger = _ProxyDBLogger() + logger = ProxyDBLogger() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): await logger.async_post_call_failure_hook( @@ -2793,19 +2793,19 @@ class TestIsMasterKey: def test_none_api_key_returns_false(self): """Regression: _is_master_key(None, 'sk-master') should return False, not raise TypeError.""" - assert _is_master_key(api_key=None, _master_key="sk-master-key") is False + assert is_master_key(api_key=None, _master_key="sk-master-key") is False def test_none_master_key_returns_false(self): - assert _is_master_key(api_key="sk-some-key", _master_key=None) is False + assert is_master_key(api_key="sk-some-key", _master_key=None) is False def test_both_none_returns_false(self): - assert _is_master_key(api_key=None, _master_key=None) is False + assert is_master_key(api_key=None, _master_key=None) is False def test_matching_key_returns_true(self): - assert _is_master_key(api_key="sk-master", _master_key="sk-master") is True + assert is_master_key(api_key="sk-master", _master_key="sk-master") is True def test_non_matching_key_returns_false(self): - assert _is_master_key(api_key="sk-other", _master_key="sk-master") is False + assert is_master_key(api_key="sk-other", _master_key="sk-master") is False def test_master_key_hash_is_rejected(self): """ @@ -2816,7 +2816,7 @@ class TestIsMasterKey: master = "sk-master-key-123" hashed = hash_token(master) - assert _is_master_key(api_key=hashed, _master_key=master) is False + assert is_master_key(api_key=hashed, _master_key=master) is False def test_sanitize_request_body_strips_secret_fields(): @@ -3199,7 +3199,7 @@ def test_sanitize_error_information_redacts_when_not_storing_prompts( ), } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert "leaked-prompt-content" not in sanitized["error_message"] @@ -3224,7 +3224,7 @@ def test_sanitize_error_information_skips_redaction_when_storing_prompts( "error_message": ('OpenAIException - {"error":{"input":[{"role":"user","content":"kept"}]}}'), } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None # User opted in via store_prompts_in_spend_logs — no key-level redaction. @@ -3251,7 +3251,7 @@ def test_sanitize_error_information_caps_size_regardless_of_prompt_flag( "error_message": huge_error, } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert len(sanitized["error_message"]) < len(huge_error) @@ -3260,7 +3260,7 @@ def test_sanitize_error_information_caps_size_regardless_of_prompt_flag( def test_sanitize_error_information_none_passthrough(): - assert _sanitize_error_information_for_spend_logs(None) is None + assert sanitize_error_information_for_spend_logs(None) is None @patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") @@ -3291,7 +3291,7 @@ def test_sanitize_error_information_reproduces_lit_2992(mock_should_store): "error_message": error_message, } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert huge_conversation_blob not in sanitized["error_message"] @@ -3370,7 +3370,7 @@ def test_sanitize_error_information_redacts_traceback_when_not_storing_prompts( "error_message": "invalid request", } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert "tb-leaked-prompt" not in sanitized["traceback"] @@ -3394,7 +3394,7 @@ def test_sanitize_error_information_skips_traceback_redaction_when_storing_promp "error_message": "invalid request", } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert "tb-kept" in sanitized["traceback"] @@ -3510,7 +3510,7 @@ def test_sanitize_error_information_redacts_pydantic_assignment_form( ), } - sanitized = _sanitize_error_information_for_spend_logs(error_info) + sanitized = sanitize_error_information_for_spend_logs(error_info) assert sanitized is not None assert "leaked-via-pydantic-msg" not in sanitized["error_message"] @@ -3535,7 +3535,7 @@ def test_sanitize_error_information_persists_no_raw_model_for_an_unknown_model_r ): error_information: Final = StandardLoggingPayloadSetup.get_error_information(original_exception=original_exception) - sanitized: Final = _sanitize_error_information_for_spend_logs( + sanitized: Final = sanitize_error_information_for_spend_logs( error_information, original_exception=original_exception ) diff --git a/tests/unit/proxy/test_aiohttp_session_recovery.py b/tests/unit/proxy/test_aiohttp_session_recovery.py index 29bd9a491b7..1ac67b76313 100644 --- a/tests/unit/proxy/test_aiohttp_session_recovery.py +++ b/tests/unit/proxy/test_aiohttp_session_recovery.py @@ -51,7 +51,7 @@ async def test_add_shared_session_recreates_closed_session(): ): with patch.object( proxy_server_module, - "_initialize_shared_aiohttp_session", + "initialize_shared_aiohttp_session", new_callable=AsyncMock, return_value=new_session, ) as mock_init: @@ -83,7 +83,7 @@ async def test_add_shared_session_handles_recreation_failure(): ): with patch.object( proxy_server_module, - "_initialize_shared_aiohttp_session", + "initialize_shared_aiohttp_session", new_callable=AsyncMock, return_value=None, ): @@ -112,7 +112,7 @@ async def test_add_shared_session_handles_recreation_exception(): ): with patch.object( proxy_server_module, - "_initialize_shared_aiohttp_session", + "initialize_shared_aiohttp_session", new_callable=AsyncMock, side_effect=RuntimeError("connection pool exhausted"), ): @@ -166,7 +166,7 @@ async def test_add_shared_session_concurrent_recreation_uses_lock(): ): with patch.object( proxy_server_module, - "_initialize_shared_aiohttp_session", + "initialize_shared_aiohttp_session", new_callable=AsyncMock, side_effect=mock_init, ): diff --git a/tests/unit/proxy/test_batch_x_litellm_model_encoding.py b/tests/unit/proxy/test_batch_x_litellm_model_encoding.py index 3161fe99e68..f8c734c4efd 100644 --- a/tests/unit/proxy/test_batch_x_litellm_model_encoding.py +++ b/tests/unit/proxy/test_batch_x_litellm_model_encoding.py @@ -91,7 +91,7 @@ async def test_create_batch_with_x_litellm_model_encodes_batch_id(): with ( patch( - "litellm.proxy.batches_endpoints.endpoints._read_request_body", + "litellm.proxy.batches_endpoints.endpoints.read_request_body", new=AsyncMock( return_value={ "input_file_id": "file-input456", @@ -206,7 +206,7 @@ async def test_create_batch_with_x_litellm_model_encodes_output_and_error_file_i with ( patch( - "litellm.proxy.batches_endpoints.endpoints._read_request_body", + "litellm.proxy.batches_endpoints.endpoints.read_request_body", new=AsyncMock( return_value={ "input_file_id": "file-input456", @@ -293,7 +293,7 @@ async def test_create_batch_without_x_litellm_model_returns_raw_ids(monkeypatch) with ( patch( - "litellm.proxy.batches_endpoints.endpoints._read_request_body", + "litellm.proxy.batches_endpoints.endpoints.read_request_body", new=AsyncMock( return_value={ "input_file_id": "file-input456", diff --git a/tests/unit/proxy/test_budget_reservation.py b/tests/unit/proxy/test_budget_reservation.py index c8e4df1030f..4635e145914 100644 --- a/tests/unit/proxy/test_budget_reservation.py +++ b/tests/unit/proxy/test_budget_reservation.py @@ -1937,7 +1937,7 @@ async def test_should_skip_reservation_when_counter_initialization_fails( return_value=0.5, ), patch( - "litellm.proxy.proxy_server._ensure_spend_counter_initialized", + "litellm.proxy.proxy_server.ensure_spend_counter_initialized", side_effect=RuntimeError("redis unavailable"), ), patch( @@ -1993,7 +1993,7 @@ async def test_should_release_tracked_entry_when_reservation_fails_after_increme side_effect=fail_after_increment, ), patch( - "litellm.proxy.proxy_server._invalidate_spend_counter", + "litellm.proxy.proxy_server.invalidate_spend_counter", side_effect=RuntimeError("invalidate unavailable"), ), ): @@ -2941,7 +2941,7 @@ async def _never_ending_stream(): def _drive_streaming_cancel(valid_token, iterator_hook): streaming_logging_obj = MagicMock() streaming_logging_obj.async_post_call_streaming_iterator_hook = iterator_hook - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect = AsyncMock() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect = AsyncMock() generator = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( response=MagicMock(), user_api_key_dict=valid_token, @@ -2985,7 +2985,7 @@ async def test_streaming_cancel_before_any_chunk_reconciles_to_input_cost( key="spend:key:key-cancel-no-chunk" ) == pytest.approx(0.5) assert reservation["finalized"] is True - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() @pytest.mark.asyncio @@ -3019,7 +3019,7 @@ async def test_streaming_cancel_after_chunk_keeps_reservation( key="spend:key:key-cancel-after-chunk" ) == pytest.approx(2.0) assert reservation.get("finalized") is not True - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() @pytest.mark.asyncio @@ -3098,7 +3098,7 @@ async def test_streaming_cancel_while_holding_back_provider_output_keeps_reserva streaming_logging_obj = MagicMock() streaming_logging_obj.async_post_call_streaming_iterator_hook = ping_then_cancel - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect = AsyncMock() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect = AsyncMock() generator = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( response=response, user_api_key_dict=valid_token, @@ -3156,7 +3156,7 @@ async def test_streaming_cancel_in_slow_path_before_yield_refunds(spend_counter_ streaming_logging_obj = MagicMock() streaming_logging_obj.async_post_call_streaming_iterator_hook = one_chunk - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect = AsyncMock() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect = AsyncMock() # On the slow path the per-chunk hook is awaited before the chunk is yielded # to the client; cancel there. Nothing has reached the client yet. streaming_logging_obj.async_post_call_streaming_hook = AsyncMock( @@ -3189,7 +3189,7 @@ async def test_streaming_cancel_in_slow_path_before_yield_refunds(spend_counter_ key="spend:key:key-cancel-slowpath" ) == pytest.approx(0.5) assert reservation["finalized"] is True - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() @pytest.mark.asyncio @@ -3219,7 +3219,7 @@ async def test_streaming_disconnect_after_consuming_chunk_keeps_reservation( key="spend:key:key-disconnect-after-chunk" ) == pytest.approx(2.0) assert reservation.get("finalized") is not True - streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + streaming_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() @pytest.mark.asyncio diff --git a/tests/unit/proxy/test_chat_completion_metadata.py b/tests/unit/proxy/test_chat_completion_metadata.py index 38dcdc13c50..7b84684d6f2 100644 --- a/tests/unit/proxy/test_chat_completion_metadata.py +++ b/tests/unit/proxy/test_chat_completion_metadata.py @@ -10,25 +10,17 @@ async def test_chat_completion_metadata_population(): # Setup request = MagicMock(spec=Request) # Mock _read_request_body to return a dict - with patch( - "litellm.proxy.proxy_server._read_request_body", new_callable=AsyncMock - ) as mock_read_body: + with patch("litellm.proxy.proxy_server.read_request_body", new_callable=AsyncMock) as mock_read_body: mock_read_body.return_value = {"model": "gpt-3.5-turbo", "messages": []} - user_api_key_dict = UserAPIKeyAuth( - user_id="test_user_id", team_id="test_team_id", org_id="test_org_id" - ) + user_api_key_dict = UserAPIKeyAuth(user_id="test_user_id", team_id="test_team_id", org_id="test_org_id") fastapi_response = MagicMock(spec=Response) # Mock ProxyBaseLLMRequestProcessing - with patch( - "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing") as MockProcessor: mock_instance = MockProcessor.return_value - mock_instance.base_process_llm_request = AsyncMock( - return_value={"choices": []} - ) + mock_instance.base_process_llm_request = AsyncMock(return_value={"choices": []}) # Execute await chat_completion( @@ -57,9 +49,7 @@ async def test_embedding_metadata_population(): from UserAPIKeyAuth. """ # Setup - with patch( - "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request" - ): + with patch("litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request"): with patch( "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.__init__", return_value=None, @@ -72,15 +62,11 @@ async def test_embedding_metadata_population(): # Create a mock Request object mock_request = MagicMock(spec=Request) - mock_request.json = AsyncMock( - return_value={"model": "gpt-3.5-turbo", "input": "hello"} - ) + mock_request.json = AsyncMock(return_value={"model": "gpt-3.5-turbo", "input": "hello"}) # Mock _read_request_body to return our data with patch( - "litellm.proxy.proxy_server._read_request_body", - new=AsyncMock( - return_value={"model": "gpt-3.5-turbo", "input": "hello"} - ), + "litellm.proxy.proxy_server.read_request_body", + new=AsyncMock(return_value={"model": "gpt-3.5-turbo", "input": "hello"}), ): # Call the endpoint function directly await embeddings( @@ -98,12 +84,8 @@ async def test_embedding_metadata_population(): else: data_arg = call_args.args[0] - assert ( - data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_emb" - ) - assert ( - data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_emb" - ) + assert data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_emb" + assert data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_emb" assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id_emb" @@ -112,28 +94,20 @@ async def test_completion_metadata_population(): # Setup request = MagicMock(spec=Request) # Mock _read_request_body to return a dict - with patch( - "litellm.proxy.proxy_server._read_request_body", new_callable=AsyncMock - ) as mock_read_body: + with patch("litellm.proxy.proxy_server.read_request_body", new_callable=AsyncMock) as mock_read_body: mock_read_body.return_value = { "model": "gpt-3.5-turbo-instruct", "prompt": "test", } - user_api_key_dict = UserAPIKeyAuth( - user_id="test_user_id_2", team_id="test_team_id_2", org_id="test_org_id_2" - ) + user_api_key_dict = UserAPIKeyAuth(user_id="test_user_id_2", team_id="test_team_id_2", org_id="test_org_id_2") fastapi_response = MagicMock(spec=Response) # Mock ProxyBaseLLMRequestProcessing - with patch( - "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing" - ) as MockProcessor: + with patch("litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing") as MockProcessor: mock_instance = MockProcessor.return_value - mock_instance.base_process_llm_request = AsyncMock( - return_value={"choices": []} - ) + mock_instance.base_process_llm_request = AsyncMock(return_value={"choices": []}) # Execute await completion( diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 368f55b1eb1..1aed60ee5e2 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -44,14 +44,14 @@ from litellm.proxy.common_request_processing import ( CostBreakdownHeaderValues, _has_attribute_error_in_chain, include_guardrail_response_requested, - _is_azure_model_router_request, + is_azure_model_router_request, open_sse_before_first_byte, resolve_litellm_call_id, ttft_keepalive_interval, _override_openai_response_model, _parse_event_data_for_error, _resolve_per_request_model_group_alias, - _should_return_raw_model_name, + should_return_raw_model_name, _sse_error_frames, _UpstreamClosingStreamingResponse, create_response, @@ -3020,7 +3020,7 @@ class TestOverrideOpenAIResponseModel: ], ) def test_raw_model_name_toggle_metadata(self, request_data, expected): - assert _should_return_raw_model_name(request_data) is expected + assert should_return_raw_model_name(request_data) is expected def test_override_model_preserves_fallback_model_when_fallback_occurred_object( self, @@ -3471,17 +3471,17 @@ class TestIsAzureModelRouterRequest: """Tests for _is_azure_model_router_request helper""" def test_detects_model_router_with_underscore(self): - assert _is_azure_model_router_request("azure_ai/model_router") is True - assert _is_azure_model_router_request("azure_ai/model_router/my-deployment") is True + assert is_azure_model_router_request("azure_ai/model_router") is True + assert is_azure_model_router_request("azure_ai/model_router/my-deployment") is True def test_detects_model_router_with_hyphen(self): - assert _is_azure_model_router_request("azure_ai/model-router") is True - assert _is_azure_model_router_request("model-router") is True + assert is_azure_model_router_request("azure_ai/model-router") is True + assert is_azure_model_router_request("model-router") is True def test_rejects_regular_models(self): - assert _is_azure_model_router_request("azure_ai/gpt-4") is False - assert _is_azure_model_router_request("gpt-4") is False - assert _is_azure_model_router_request("openai/gpt-3.5-turbo") is False + assert is_azure_model_router_request("azure_ai/gpt-4") is False + assert is_azure_model_router_request("gpt-4") is False + assert is_azure_model_router_request("openai/gpt-3.5-turbo") is False class TestStreamingOverheadHeader: @@ -5263,7 +5263,7 @@ class TestStreamingClientDisconnectLogging: fire_spy = MagicMock() monkeypatch.setattr( - "litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging", + "litellm.proxy.utils.ProxyLogging.fire_deferred_stream_logging", fire_spy, ) @@ -5297,7 +5297,7 @@ class TestStreamingClientDisconnectLogging: fire_spy = MagicMock() monkeypatch.setattr( - "litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging", + "litellm.proxy.utils.ProxyLogging.fire_deferred_stream_logging", fire_spy, ) @@ -5328,7 +5328,7 @@ class TestStreamingClientDisconnectLogging: ) monkeypatch.setattr( - "litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging", + "litellm.proxy.utils.ProxyLogging.fire_deferred_stream_logging", MagicMock(), ) @@ -6944,7 +6944,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: ProxyRateLimitError, ) from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache @@ -6962,7 +6962,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: # Real per-key per-model TPM limiter + a key carrying the customer's # `model_tpm_limit` metadata (only the primary is capped). - limiter = _PROXY_MaxParallelRequestsHandler( + limiter = PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(DualCache()) ) user_api_key_dict = UserAPIKeyAuth( @@ -7064,11 +7064,11 @@ class TestPreCallWithFallbacksOnLocalRateLimit: ``add_litellm_data_to_request`` with a live OTel span, ``function_setup``, then the limiter.""" from litellm.caching.caching import DualCache from litellm.proxy import proxy_server - from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.hooks.parallel_request_limiter_v3 import PROXY_MaxParallelRequestsHandler_v3 from litellm.proxy.utils import InternalUsageCache monkeypatch.setattr(proxy_server, "prisma_client", None) - limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) + limiter = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) limiter_models: list[str] = [] async def run_limiter( @@ -7561,7 +7561,7 @@ class TestStreamingClientDisconnectBilling: try: response = await self._start_partial_stream() proxy_logging_obj = types.SimpleNamespace( - _arelease_max_parallel_requests_on_disconnect=AsyncMock(), + arelease_max_parallel_requests_on_disconnect=AsyncMock(), ) billed = await _bill_partial_streamed_spend_on_disconnect( @@ -7581,7 +7581,7 @@ class TestStreamingClientDisconnectBilling: finally: litellm.callbacks = original_callbacks - proxy_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_not_called() + proxy_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_not_called() @pytest.mark.asyncio async def test_disconnect_without_billable_chunks_releases_slot(self): @@ -7596,7 +7596,7 @@ class TestStreamingClientDisconnectBilling: # No chunks to assemble -> billing dispatches no success event. empty_response = types.SimpleNamespace(chunks=[], messages=None) proxy_logging_obj = types.SimpleNamespace( - _arelease_max_parallel_requests_on_disconnect=AsyncMock(), + arelease_max_parallel_requests_on_disconnect=AsyncMock(), ) await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( @@ -7609,7 +7609,7 @@ class TestStreamingClientDisconnectBilling: proxy_logging_obj=proxy_logging_obj, ) - proxy_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() + proxy_logging_obj.arelease_max_parallel_requests_on_disconnect.assert_awaited_once() async def _bill_and_collect_success_event(self, prepare=None, request_data=None): recorder = _RecordingSuccessLogger() @@ -8287,7 +8287,7 @@ class TestPerRequestModelGroupAlias: monkeypatch.setattr( litellm.proxy.common_request_processing, - "_check_and_merge_model_level_guardrails", + "check_and_merge_model_level_guardrails", recording_merge, ) diff --git a/tests/unit/proxy/test_dynamic_mcp_route.py b/tests/unit/proxy/test_dynamic_mcp_route.py index 83963fbd962..4a8dddb1bf4 100644 --- a/tests/unit/proxy/test_dynamic_mcp_route.py +++ b/tests/unit/proxy/test_dynamic_mcp_route.py @@ -33,7 +33,7 @@ _IS_ACCESS_GROUP = "litellm.proxy.proxy_server._is_mcp_access_group_cached" _USER_API_KEY_CACHE = "litellm.proxy.proxy_server.user_api_key_cache" _GET_ACCESS_GROUP_SERVERS = ( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." - "MCPRequestHandler._get_mcp_servers_from_access_groups" + "MCPRequestHandler.get_mcp_servers_from_access_groups" ) _FORWARD = "litellm.proxy.proxy_server._mcp_forward_as_path" _RESOLVE_CSV = "litellm.proxy.proxy_server._resolve_mcp_csv_tokens" @@ -286,10 +286,10 @@ async def test_dynamic_mcp_route_resolves_toolset(): async def fake_stream(fn, scope, receive): nonlocal captured_toolset_id from litellm.proxy._experimental.mcp_server.server import ( - _mcp_active_toolset_id, + mcp_active_toolset_id, ) - captured_toolset_id = _mcp_active_toolset_id.get() + captured_toolset_id = mcp_active_toolset_id.get() captured_scope.update(scope) with ( diff --git a/tests/unit/proxy/test_health_check_functions.py b/tests/unit/proxy/test_health_check_functions.py index 1b2fc73fca7..2c0198cd538 100644 --- a/tests/unit/proxy/test_health_check_functions.py +++ b/tests/unit/proxy/test_health_check_functions.py @@ -12,7 +12,7 @@ from litellm.proxy.health_endpoints._health_endpoints import ( _aggregate_health_check_results, _build_model_param_to_info_mapping, _perform_health_check_and_save, - _save_background_health_checks_to_db, + save_background_health_checks_to_db, _save_health_check_results_if_changed, _save_health_check_to_db, latest_health_checks_endpoint, @@ -427,7 +427,7 @@ async def test_save_background_health_checks_to_db(): start_time = 1234567890.0 - persisted = await _save_background_health_checks_to_db( + persisted = await save_background_health_checks_to_db( mock_prisma, model_list, healthy_endpoints, @@ -535,7 +535,7 @@ async def test_save_background_health_checks_to_db_returns_false_when_a_write_fa mock_prisma.save_health_check_result = AsyncMock(return_value=None) model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() - persisted = await _save_background_health_checks_to_db( + persisted = await save_background_health_checks_to_db( mock_prisma, model_list, healthy_endpoints, unhealthy_endpoints, 1234567890.0, "background_health_check" ) @@ -552,7 +552,7 @@ async def test_save_background_health_checks_to_db_writes_nothing_when_the_lates mock_prisma.save_health_check_result = AsyncMock(return_value={"id": "row"}) model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() - persisted = await _save_background_health_checks_to_db( + persisted = await save_background_health_checks_to_db( mock_prisma, model_list, healthy_endpoints, unhealthy_endpoints, 1234567890.0, "background_health_check" ) @@ -562,7 +562,7 @@ async def test_save_background_health_checks_to_db_writes_nothing_when_the_lates @pytest.mark.asyncio async def test_save_background_health_checks_to_db_no_prisma(): """Test graceful handling when no prisma client""" - result = await _save_background_health_checks_to_db(None, [], [], [], 0.0, "background_health_check") + result = await save_background_health_checks_to_db(None, [], [], [], 0.0, "background_health_check") assert result is False @@ -582,7 +582,7 @@ async def test_save_background_health_checks_to_db_exception_handling(): # Must not raise (the health check loop has to survive a DB outage) but must report # the failure, so the window lock can be released for another pod to retry - persisted = await _save_background_health_checks_to_db( + persisted = await save_background_health_checks_to_db( mock_prisma, model_list, [], [], 0.0, "background_health_check" ) @@ -653,7 +653,7 @@ async def test_save_background_health_checks_compares_raw_checked_at_against_utc {"model_name": "fresh-model", "model_info": {"id": "fresh-id"}, "litellm_params": {"model": "openai/fresh"}}, ] - await _save_background_health_checks_to_db( + await save_background_health_checks_to_db( mock_prisma, model_list, [{"model": "openai/stale"}, {"model": "openai/fresh"}], diff --git a/tests/unit/proxy/test_health_check_max_tokens.py b/tests/unit/proxy/test_health_check_max_tokens.py index e3641ac2c81..091de1e24b3 100644 --- a/tests/unit/proxy/test_health_check_max_tokens.py +++ b/tests/unit/proxy/test_health_check_max_tokens.py @@ -12,7 +12,7 @@ from litellm.proxy.health_check import ( _is_strategy_router_deployment, _resolve_health_check_max_tokens, resolve_health_check_mode, - _update_litellm_params_for_health_check, + update_litellm_params_for_health_check, ) @@ -26,7 +26,7 @@ async def test_update_litellm_params_max_tokens_default(monkeypatch): model_info = {} litellm_params = {"model": "gpt-4"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 16 @@ -39,7 +39,7 @@ async def test_update_litellm_params_max_tokens_custom(): model_info = {"health_check_max_tokens": 5} litellm_params = {"model": "gpt-4"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 5 @@ -52,7 +52,7 @@ async def test_update_litellm_params_max_tokens_wildcard(): model_info = {} litellm_params = {"model": "openai/*"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert "max_tokens" not in updated_params @@ -102,7 +102,7 @@ async def test_background_health_check_max_tokens_env_var(monkeypatch): model_info = {} litellm_params = {"model": "azure/gpt-4"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 10 @@ -118,7 +118,7 @@ async def test_per_model_overrides_global_env_var(monkeypatch): model_info = {"health_check_max_tokens": 5} litellm_params = {"model": "azure/gpt-4"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 5 @@ -133,7 +133,7 @@ async def test_global_env_var_applies_to_wildcard_models(monkeypatch): model_info = {} litellm_params = {"model": "openai/*"} - updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + updated_params = update_litellm_params_for_health_check(model_info, litellm_params) assert updated_params["max_tokens"] == 15 @@ -184,12 +184,12 @@ async def test_background_split_env_reasoning_vs_non_reasoning(monkeypatch): litellm_params = {"model": "azure/gpt-4"} with patch.object(hc_module.litellm, "supports_reasoning", return_value=False): - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 16 litellm_params2 = {"model": "openai/o1"} with patch.object(hc_module.litellm, "supports_reasoning", return_value=True): - updated2 = _update_litellm_params_for_health_check(model_info, litellm_params2) + updated2 = update_litellm_params_for_health_check(model_info, litellm_params2) assert updated2["max_tokens"] == 50 @@ -202,7 +202,7 @@ async def test_reasoning_env_precedence_over_global(monkeypatch): litellm_params = {"model": "openai/gpt-5.4"} with patch.object(hc_module.litellm, "supports_reasoning", return_value=True): - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 20 @@ -215,7 +215,7 @@ async def test_non_reasoning_uses_global_when_reasoning_env_set(monkeypatch): litellm_params = {"model": "azure/gpt-4"} with patch.object(hc_module.litellm, "supports_reasoning", return_value=False): - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 10 @@ -250,7 +250,7 @@ def test_image_generation_mode_skips_max_tokens(): model_info = {"mode": "image_generation"} litellm_params = {"model": "openai/dall-e-3", "api_key": "sk-test"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert "max_tokens" not in updated # connection-level params must still pass through unchanged @@ -267,7 +267,7 @@ def test_health_check_max_tokens_value_is_ignored_for_non_chat_modes(): model_info = {"mode": "image_generation", "health_check_max_tokens": 50} litellm_params = {"model": "openai/dall-e-3"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert "max_tokens" not in updated @@ -277,7 +277,7 @@ def test_chat_mode_still_injects_max_tokens(): model_info = {"mode": "chat"} litellm_params = {"model": "gpt-4"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 16 @@ -287,7 +287,7 @@ def test_no_mode_still_injects_max_tokens(): model_info: dict = {} litellm_params = {"model": "gpt-4"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 16 @@ -305,7 +305,7 @@ def test_no_mode_still_injects_max_tokens(): @pytest.mark.parametrize("mode", ["chat", "completion", "responses"]) def test_chat_style_modes_inject_max_tokens(mode): - updated = _update_litellm_params_for_health_check({"mode": mode}, {"model": f"openai/dummy-{mode}"}) + updated = update_litellm_params_for_health_check({"mode": mode}, {"model": f"openai/dummy-{mode}"}) assert updated["max_tokens"] == 16 @@ -326,7 +326,7 @@ def test_chat_style_modes_inject_max_tokens(mode): ], ) def test_non_chat_modes_skip_max_tokens(mode): - updated = _update_litellm_params_for_health_check({"mode": mode}, {"model": f"openai/dummy-{mode}"}) + updated = update_litellm_params_for_health_check({"mode": mode}, {"model": f"openai/dummy-{mode}"}) assert "max_tokens" not in updated @@ -339,7 +339,7 @@ def test_explicit_override_true_forces_injection_outside_allowlist(): } litellm_params = {"model": "openai/some-future-image-model"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert updated["max_tokens"] == 16 @@ -349,7 +349,7 @@ def test_explicit_override_false_suppresses_injection_inside_allowlist(): model_info = {"mode": "chat", "health_check_supports_max_tokens": False} litellm_params = {"model": "openai/strict-schema-chat"} - updated = _update_litellm_params_for_health_check(model_info, litellm_params) + updated = update_litellm_params_for_health_check(model_info, litellm_params) assert "max_tokens" not in updated @@ -358,29 +358,29 @@ def test_update_litellm_params_health_check_reasoning_effort(): """model_info.health_check_reasoning_effort sets reasoning_effort for chat-style health checks.""" model_info = {"health_check_reasoning_effort": "low"} litellm_params = {"model": "openai/gpt-5", "api_key": "x"} - out = _update_litellm_params_for_health_check(model_info, dict(litellm_params)) + out = update_litellm_params_for_health_check(model_info, dict(litellm_params)) assert out.get("reasoning_effort") == "low" model_info = {"mode": "chat", "health_check_reasoning_effort": "none"} - out = _update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5", "api_key": "x"}) assert out.get("reasoning_effort") == "none" model_info = {"mode": "completion", "health_check_reasoning_effort": "low"} - out = _update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5", "api_key": "x"}) assert out.get("reasoning_effort") == "low" model_info = { "health_check_reasoning_effort": {"effort": "none", "summary": "auto"}, } - out = _update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5.1", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-5.1", "api_key": "x"}) assert out.get("reasoning_effort") == {"effort": "none", "summary": "auto"} model_info = {"mode": "embedding", "health_check_reasoning_effort": "low"} - out = _update_litellm_params_for_health_check(model_info, {"model": "text-embedding-3-small", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "text-embedding-3-small", "api_key": "x"}) assert "reasoning_effort" not in out model_info = {} - out = _update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-4o", "api_key": "x"}) + out = update_litellm_params_for_health_check(model_info, {"model": "openai/gpt-4o", "api_key": "x"}) assert "reasoning_effort" not in out @@ -408,7 +408,7 @@ def test_bedrock_embedding_without_explicit_mode_skips_max_tokens(deployment_mod """Embedding mode auto-detected from model cost map -> no max_tokens, provider pinned.""" assert resolve_health_check_mode({}, {"model": deployment_model}) == "embedding" - updated = _update_litellm_params_for_health_check({}, {"model": deployment_model}) + updated = update_litellm_params_for_health_check({}, {"model": deployment_model}) assert "max_tokens" not in updated assert updated["custom_llm_provider"] == "bedrock" @@ -427,7 +427,7 @@ def test_resolve_health_check_mode_unknown_model_returns_none(): def test_bedrock_chat_without_mode_still_injects_max_tokens_and_pins_provider(): """Regression guard: chat-style Bedrock deployments keep max_tokens and get the provider pin.""" - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {}, {"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"} ) @@ -443,7 +443,7 @@ def test_bedrock_prefix_strip_preserves_explicit_custom_llm_provider(): not clobber a more specific one, otherwise a converse deployment would be probed against the Invoke endpoint and report a spurious failure. """ - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {}, { "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", @@ -509,7 +509,7 @@ def test_mantle_claude_without_mode_resolves_to_anthropic_messages(deployment_mo """Mantle only serves Claude over /anthropic/v1/messages, so that is the probe surface by default.""" assert resolve_health_check_mode({}, {"model": deployment_model}) == "anthropic_messages" - updated = _update_litellm_params_for_health_check({}, {"model": deployment_model}) + updated = update_litellm_params_for_health_check({}, {"model": deployment_model}) assert updated["max_tokens"] == 16 assert [message["role"] for message in updated["messages"]] == ["user"] @@ -597,7 +597,7 @@ def test_autodetected_embedding_skips_reasoning_effort(): Bedrock embedding probe, which embeddings reject as an unknown field. The mode is now resolved from the cost map, so embeddings are excluded. """ - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {"health_check_reasoning_effort": "low"}, {"model": "bedrock/amazon.titan-embed-text-v2:0"}, ) @@ -649,7 +649,7 @@ def test_health_check_params_merge_into_probe_params(): """health_check_params reach the probe request for the deployment that declares them.""" media_source = {"s3Location": {"uri": "s3://my-bucket/clip.mp4"}} - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {"mode": "chat", "health_check_params": {"mediaSource": media_source}}, {"model": "bedrock/us.twelvelabs.pegasus-1-2-v1:0"}, ) @@ -674,7 +674,7 @@ def test_health_check_params_lose_to_dedicated_health_check_knobs(): "health_check_reasoning_effort": "none", } - updated = _update_litellm_params_for_health_check(model_info, {"model": "openai/dummy"}) + updated = update_litellm_params_for_health_check(model_info, {"model": "openai/dummy"}) assert updated["max_tokens"] == 5 assert updated["model"] == "openai/cheap-model" @@ -684,7 +684,7 @@ def test_health_check_params_lose_to_dedicated_health_check_knobs(): def test_health_check_params_lose_to_the_audio_speech_voice_knob(): """health_check_voice still wins for audio_speech deployments.""" - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( { "mode": "audio_speech", "health_check_params": {"voice": "sage", "response_format": "wav"}, @@ -704,7 +704,7 @@ def test_health_check_params_lose_to_the_audio_speech_voice_knob(): def test_health_check_params_ignored_when_not_a_dict(bad_value, caplog): """A misconfigured health_check_params is skipped with a warning instead of breaking the probe.""" with caplog.at_level(logging.WARNING, logger="litellm.proxy.health_check"): - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {"mode": "chat", "health_check_params": bad_value}, {"model": "openai/dummy"}, ) @@ -716,7 +716,7 @@ def test_health_check_params_ignored_when_not_a_dict(bad_value, caplog): def test_health_check_params_apply_to_non_chat_modes(): """Non-chat probes get health_check_params too, and still no max_tokens.""" - updated = _update_litellm_params_for_health_check( + updated = update_litellm_params_for_health_check( {"mode": "embedding", "health_check_params": {"dimensions": 8}}, {"model": "bedrock/amazon.titan-embed-text-v2:0"}, ) @@ -731,7 +731,7 @@ async def _pegasus_health_check_request_body( monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) litellm.in_memory_llm_clients_cache.flush_cache() - litellm_params = _update_litellm_params_for_health_check( + litellm_params = update_litellm_params_for_health_check( model_info, { "model": "bedrock/us.twelvelabs.pegasus-1-2-v1:0", diff --git a/tests/unit/proxy/test_litellm_pre_call_utils.py b/tests/unit/proxy/test_litellm_pre_call_utils.py index d4624719565..86492537fd3 100644 --- a/tests/unit/proxy/test_litellm_pre_call_utils.py +++ b/tests/unit/proxy/test_litellm_pre_call_utils.py @@ -22,9 +22,9 @@ from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _apply_credential_overrides_from_model_config, _extract_credential_from_entry, - _get_dynamic_logging_metadata, + get_dynamic_logging_metadata, _get_enforced_params, - _get_metadata_variable_name, + get_metadata_variable_name, _match_and_track_policies, _promoted_trace_control_fields, _resolve_credential_from_model_config, @@ -84,45 +84,45 @@ class TestGetMetadataVariableName: def test_returns_litellm_metadata_for_thread_routes(self): request = self._make_request("/v1/threads/thread_123/messages") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_assistant_routes(self): request = self._make_request("/v1/assistants/asst_123") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_batches_route(self): request = self._make_request("/v1/batches") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_messages_route(self): request = self._make_request("/v1/messages") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_files_route(self): request = self._make_request("/v1/files") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_metadata_for_chat_completions(self): request = self._make_request("/chat/completions") - assert _get_metadata_variable_name(request) == "metadata" + assert get_metadata_variable_name(request) == "metadata" def test_returns_metadata_for_completions(self): request = self._make_request("/v1/completions") - assert _get_metadata_variable_name(request) == "metadata" + assert get_metadata_variable_name(request) == "metadata" def test_returns_metadata_for_embeddings(self): request = self._make_request("/v1/embeddings") - assert _get_metadata_variable_name(request) == "metadata" + assert get_metadata_variable_name(request) == "metadata" def test_returns_litellm_metadata_for_bedrock_invoke(self): # GH#30629: bedrock passthrough must use litellm_metadata # to prevent key-level tags from leaking into provider body request = self._make_request("/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_returns_litellm_metadata_for_bedrock_converse(self): request = self._make_request("/bedrock/model/us.anthropic.claude-sonnet-4-6/converse") - assert _get_metadata_variable_name(request) == "litellm_metadata" + assert get_metadata_variable_name(request) == "litellm_metadata" def test_get_enforced_params_for_service_account_settings(): @@ -902,7 +902,7 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l monkeypatch: pytest.MonkeyPatch, pre_call_ran: bool ) -> None: from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.proxy.guardrails.guardrail_hooks.presidio import _OPTIONAL_PresidioPIIMasking + from litellm.proxy.guardrails.guardrail_hooks.presidio import OPTIONAL_PresidioPIIMasking from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot from litellm.proxy.spend_tracking.spend_tracking_utils import _get_proxy_server_request_for_spend_logs_payload @@ -921,7 +921,7 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l logging_obj.update_messages(messages) snapshot: Final = logging_obj.shadow_eval_request_snapshot assert (snapshot is not None) is pre_call_ran - guardrail: Final = _OPTIONAL_PresidioPIIMasking( + guardrail: Final = OPTIONAL_PresidioPIIMasking( mock_testing=True, logging_only=True, mock_redacted_text={"text": "email [EMAIL]", "items": []} ) @@ -2419,7 +2419,7 @@ def test_get_dynamic_logging_metadata_with_arize_team_logging(): mock_proxy_config = MagicMock() # Call the function - result = _get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=mock_proxy_config) + result = get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=mock_proxy_config) # Verify the result assert result is not None @@ -2466,7 +2466,7 @@ def test_get_dynamic_logging_metadata_ignores_env_reference_from_key_metadata( team_metadata={}, ) - result = _get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=MagicMock()) + result = get_dynamic_logging_metadata(user_api_key_dict=user_api_key_dict, proxy_config=MagicMock()) assert result is None @@ -4008,7 +4008,7 @@ async def test_team_guardrails_append_to_key_guardrails(): team_metadata={"guardrails": ["team-guardrail-1", "key-guardrail-1"]}, ) - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): updated_data = await add_litellm_data_to_request( data=data, request=request_mock, @@ -4057,7 +4057,7 @@ async def test_request_guardrails_do_not_override_key_guardrails(): "guardrails": [], } - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): updated_data_empty = await add_litellm_data_to_request( data=data_with_empty, request=request_mock, @@ -4103,7 +4103,7 @@ async def test_project_guardrails_merge_with_key_and_team(): project_metadata={"guardrails": ["project-guardrail-1", "team-guardrail-1"]}, ) - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): updated_data = await add_litellm_data_to_request( data=data, request=request_mock, @@ -4152,7 +4152,7 @@ async def test_project_guardrails_only(): project_metadata={"guardrails": ["project-guardrail-1", "project-guardrail-2"]}, ) - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): updated_data = await add_litellm_data_to_request( data=data, request=request_mock, @@ -5459,7 +5459,7 @@ async def test_team_guardrail_merges_with_global_policy(): attachment_registry._initialized = True try: - with patch("litellm.proxy.utils._premium_user_check"): + with patch("litellm.proxy.utils.premium_user_check"): await move_guardrails_to_metadata( data=data, _metadata_variable_name="metadata", @@ -7127,7 +7127,7 @@ class TestPromotedTraceControlFields: ) def test_returns_litellm_metadata_for_responses_route(self): - assert _get_metadata_variable_name(self._make_request("/v1/responses")) == "litellm_metadata" + assert get_metadata_variable_name(self._make_request("/v1/responses")) == "litellm_metadata" def test_promotes_trace_prefixed_and_allow_listed_fields(self): requester_metadata = { diff --git a/tests/unit/proxy/test_model_level_guardrails.py b/tests/unit/proxy/test_model_level_guardrails.py index 9eaae49c46a..0bc9c348127 100644 --- a/tests/unit/proxy/test_model_level_guardrails.py +++ b/tests/unit/proxy/test_model_level_guardrails.py @@ -15,7 +15,7 @@ from unittest.mock import AsyncMock, MagicMock, patch sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))) from litellm.proxy.utils import ( - _check_and_merge_model_level_guardrails, + check_and_merge_model_level_guardrails, _merge_guardrails_with_existing, ) @@ -38,7 +38,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = ["openai-moderation"] mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert "openai-moderation" in result["metadata"]["guardrails"] mock_router.get_deployment.assert_called_once_with(model_id="model-uuid-123") @@ -57,7 +57,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = ["model-guardrail"] mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert "existing-guardrail" in result["metadata"]["guardrails"] assert "model-guardrail" in result["metadata"]["guardrails"] @@ -76,14 +76,14 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = ["openai-moderation"] mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert result["metadata"]["guardrails"].count("openai-moderation") == 1 def test_returns_data_unchanged_when_no_router(self): """Returns data unchanged when llm_router is None.""" data = {"model": "gpt-4", "metadata": {}} - result = _check_and_merge_model_level_guardrails(data=data, llm_router=None) + result = check_and_merge_model_level_guardrails(data=data, llm_router=None) assert result is data def test_returns_data_unchanged_when_no_model_info(self): @@ -95,7 +95,7 @@ class TestCheckAndMergeModelLevelGuardrails: # finds a deployment. mock_router.get_deployment.return_value = None mock_router.get_deployment_by_model_group_name.return_value = None - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert result is data def test_returns_data_unchanged_when_deployment_has_no_guardrails(self): @@ -109,7 +109,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = None mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert result is data @@ -122,7 +122,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_router = MagicMock() mock_router.get_deployment.return_value = None - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) assert result is data @@ -140,7 +140,7 @@ class TestCheckAndMergeModelLevelGuardrails: mock_deployment.litellm_params.get.return_value = ["new-guardrail"] mock_router.get_deployment.return_value = mock_deployment - result = _check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) + result = check_and_merge_model_level_guardrails(data=data, llm_router=mock_router) # Result is a different top-level dict assert result is not data diff --git a/tests/unit/proxy/test_model_list_callback_filter.py b/tests/unit/proxy/test_model_list_callback_filter.py index 00fbfee24ed..d397ee676f2 100644 --- a/tests/unit/proxy/test_model_list_callback_filter.py +++ b/tests/unit/proxy/test_model_list_callback_filter.py @@ -107,7 +107,7 @@ def team_admin_privileges(monkeypatch) -> None: async def _is_team_admin(**kwargs) -> bool: return True - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _is_team_admin) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _is_team_admin) def _non_admin(**kwargs) -> UserAPIKeyAuth: diff --git a/tests/unit/proxy/test_model_list_discoverable.py b/tests/unit/proxy/test_model_list_discoverable.py index bcd52479f2c..0be5342c45c 100644 --- a/tests/unit/proxy/test_model_list_discoverable.py +++ b/tests/unit/proxy/test_model_list_discoverable.py @@ -71,7 +71,7 @@ def team_admin_privileges(monkeypatch) -> None: async def _is_team_admin(**kwargs) -> bool: return True - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _is_team_admin) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _is_team_admin) def _non_admin() -> UserAPIKeyAuth: diff --git a/tests/unit/proxy/test_model_list_healthy_only.py b/tests/unit/proxy/test_model_list_healthy_only.py index 718c7e41da8..2652fbb30ef 100644 --- a/tests/unit/proxy/test_model_list_healthy_only.py +++ b/tests/unit/proxy/test_model_list_healthy_only.py @@ -120,7 +120,7 @@ async def test_model_list_healthy_only_applies_to_scope_expand( async def _fake_admin(**kwargs): return True - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _fake_admin) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _fake_admin) monkeypatch.setattr( model_checks, "get_complete_model_list", @@ -158,7 +158,7 @@ async def test_model_list_general_setting_applies_to_scope_expand(patched_model_ async def _fake_admin(**kwargs): return True - monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _fake_admin) + monkeypatch.setattr(common_utils, "user_has_admin_privileges", _fake_admin) monkeypatch.setattr( model_checks, "get_complete_model_list", diff --git a/tests/unit/proxy/test_modify_response_streaming_passthrough.py b/tests/unit/proxy/test_modify_response_streaming_passthrough.py index da57d9c616e..f4d3b659ee1 100644 --- a/tests/unit/proxy/test_modify_response_streaming_passthrough.py +++ b/tests/unit/proxy/test_modify_response_streaming_passthrough.py @@ -40,7 +40,7 @@ async def _run_streaming_block_and_get_wrapper(exception): outer_body = {"model": "gpt-4o", "messages": [], "stream": True} with patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", new_callable=AsyncMock, return_value=outer_body, ), patch( diff --git a/tests/unit/proxy/test_proxy_cli.py b/tests/unit/proxy/test_proxy_cli.py index 8fc807cd8a0..c73b60f3c26 100644 --- a/tests/unit/proxy/test_proxy_cli.py +++ b/tests/unit/proxy/test_proxy_cli.py @@ -607,7 +607,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -713,7 +713,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.side_effect = lambda *a, **k: { @@ -806,7 +806,7 @@ class TestProxyInitializationHelpers: }, ), patch( # test-quality-ok: same isolation as the sibling CLI tests above - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.side_effect = lambda *a, **k: { @@ -935,7 +935,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1065,7 +1065,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1186,7 +1186,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1279,7 +1279,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1358,7 +1358,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( "litellm.proxy.proxy_cli.append_query_params", @@ -1417,7 +1417,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -1478,10 +1478,10 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._is_port_in_use", + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.is_port_in_use", return_value=False, ), ): @@ -1553,10 +1553,10 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._is_port_in_use", + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.is_port_in_use", return_value=False, ), ): @@ -1620,7 +1620,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -1678,7 +1678,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -1706,7 +1706,7 @@ class TestProxyInitializationHelpers: assert call_args[1]["limit_max_requests"] == 1000 assert call_args[1]["limit_max_requests_jitter"] == 50 - @patch("litellm.proxy.proxy_cli.ProxyInitializationHelpers._run_gunicorn_server") + @patch("litellm.proxy.proxy_cli.ProxyInitializationHelpers.run_gunicorn_server") @patch("uvicorn.run") @patch("builtins.print") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") @@ -1738,7 +1738,7 @@ class TestProxyInitializationHelpers: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2015,7 +2015,7 @@ class TestProxyInitializationHelpers: {"litellm.proxy.proxy_server": mock_proxy_server_module}, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2080,7 +2080,7 @@ class TestQueryEngineReaperWiring: "litellm.proxy.proxy_cli.start_query_engine_reaper" ) as mock_start_reaper, patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2180,7 +2180,7 @@ class TestRunServerDbSetup: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2257,7 +2257,7 @@ class TestRunServerDbSetup: }, ), patch( # test-quality-ok: same isolation as the sibling CLI tests above - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2374,7 +2374,7 @@ class TestRunServerDbSetup: }, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { @@ -2743,7 +2743,7 @@ class TestRunServerDbSetup: {"proxy_server": mock_proxy_module, "litellm.proxy.proxy_server": mock_proxy_module}, ), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, outcome as exc_info, ): @@ -3435,7 +3435,7 @@ class TestTokenAuthCliFlags: patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database"), patch("uvicorn.run"), patch( - "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + "litellm.proxy.proxy_cli.ProxyInitializationHelpers.get_default_unvicorn_init_args" ) as mock_get_args, ): mock_get_args.return_value = { diff --git a/tests/unit/proxy/test_proxy_server.py b/tests/unit/proxy/test_proxy_server.py index 9772988eac8..79f4edd4d62 100644 --- a/tests/unit/proxy/test_proxy_server.py +++ b/tests/unit/proxy/test_proxy_server.py @@ -1238,7 +1238,7 @@ async def test_team_update_redis(): """ from litellm.caching.caching import DualCache, RedisCache from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.auth_checks import _cache_team_object + from litellm.proxy.auth.auth_checks import cache_team_object proxy_logging_obj: ProxyLogging = getattr( litellm.proxy.proxy_server, "proxy_logging_obj" @@ -1251,7 +1251,7 @@ async def test_team_update_redis(): "async_set_cache", new=AsyncMock(), ) as mock_client: - await _cache_team_object( + await cache_team_object( team_id="1234", team_table=LiteLLM_TeamTableCachedObj(team_id="1234"), user_api_key_cache=DualCache(redis_cache=redis_cache), 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 ec91f1d8edc..920f9d90bb8 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -1652,7 +1652,7 @@ mock_prisma = MockPrisma() @patch( - "litellm.proxy.proxy_server.ProxyStartupEvent._setup_prisma_client", + "litellm.proxy.proxy_server.ProxyStartupEvent.setup_prisma_client", return_value=mock_prisma, ) @pytest.mark.asyncio @@ -3623,12 +3623,12 @@ async def test_startup_initializes_string_callbacks_after_all_litellm_settings_l def test_startup_hands_router_to_every_registered_prompt_injection_detector(monkeypatch): from litellm.proxy._types import LiteLLMPromptInjectionParams - from litellm.proxy.hooks.prompt_injection_detection import _OPTIONAL_PromptInjectionDetection + from litellm.proxy.hooks.prompt_injection_detection import OPTIONAL_PromptInjectionDetection from litellm.proxy.proxy_server import ProxyStartupEvent from litellm.router import Router monkeypatch.setattr(litellm, "callbacks", []) - detector = _OPTIONAL_PromptInjectionDetection( + detector = OPTIONAL_PromptInjectionDetection( prompt_injection_params=LiteLLMPromptInjectionParams( heuristics_check=False, llm_api_check=True, @@ -4545,7 +4545,7 @@ async def test_chat_completion_result_no_nested_none_values(): with ( patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", return_value={"model": "gpt-3.5-turbo", "messages": []}, ), patch( @@ -6554,7 +6554,7 @@ async def test_init_sso_settings_in_db(): mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config) # Mock _decrypt_and_set_db_env_variables - with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set: + with patch.object(proxy_config, "decrypt_and_set_db_env_variables") as mock_decrypt_and_set: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) # Verify find_unique was called with correct parameters @@ -6598,7 +6598,7 @@ async def test_init_sso_settings_in_db_no_settings(): mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) # Mock _decrypt_and_set_db_env_variables - with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set: + with patch.object(proxy_config, "decrypt_and_set_db_env_variables") as mock_decrypt_and_set: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) # Verify find_unique was called @@ -6652,7 +6652,7 @@ async def test_init_sso_settings_in_db_empty_settings(): mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config) # Mock _decrypt_and_set_db_env_variables - with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set: + with patch.object(proxy_config, "decrypt_and_set_db_env_variables") as mock_decrypt_and_set: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) # Verify find_unique was called @@ -6694,7 +6694,7 @@ async def test_init_sso_settings_in_db_retries_on_transport_error(): mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 - with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt: + with patch.object(proxy_config, "decrypt_and_set_db_env_variables") as mock_decrypt: await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) assert len(invocations) == 2 @@ -6872,7 +6872,7 @@ def test_root_redirect_when_docs_url_not_root_and_redirect_url_set(monkeypatch): from fastapi.responses import RedirectResponse from litellm.proxy.proxy_server import cleanup_router_config_variables - from litellm.proxy.utils import _get_docs_url + from litellm.proxy.utils import get_docs_url cleanup_router_config_variables() filepath = os.path.dirname(os.path.abspath(__file__)) @@ -6885,7 +6885,7 @@ def test_root_redirect_when_docs_url_not_root_and_redirect_url_set(monkeypatch): asyncio.run(initialize(config=config_fp, debug=True)) - docs_url = _get_docs_url() + docs_url = get_docs_url() root_redirect_url = os.getenv("ROOT_REDIRECT_URL") # Remove any existing "/" route that might interfere @@ -7368,7 +7368,7 @@ class TestInvitationEndpoints: mock_prisma.db.litellm_invitationlink = MagicMock() # Avoid triggering async DB calls in _user_has_admin_privileges with patch( - "litellm.proxy.proxy_server._user_has_admin_privileges", + "litellm.proxy.proxy_server.user_has_admin_privileges", new_callable=AsyncMock, return_value=False, ): @@ -7537,7 +7537,7 @@ async def test_async_data_generator_uses_direct_stream_fast_path_without_callbac mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging") as mock_deferred_logging: + with patch.object(ProxyLogging, "fire_deferred_stream_logging") as mock_deferred_logging: yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -7593,7 +7593,7 @@ async def test_async_data_generator_preserves_non_raw_sse_like_bytes(): mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -7650,7 +7650,7 @@ async def test_async_data_generator_buffers_split_google_native_sse_json_frame() mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -7698,7 +7698,7 @@ async def test_async_data_generator_flushes_raw_sse_stream_without_trailing_deli with ( patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), - patch.object(ProxyLogging, "_fire_deferred_stream_logging"), + patch.object(ProxyLogging, "fire_deferred_stream_logging"), ): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): @@ -7748,7 +7748,7 @@ async def test_async_data_generator_errors_when_raw_sse_frame_exceeds_buffer_lim with ( patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), patch("litellm.proxy.proxy_server._MAX_RAW_SSE_BUFFER_CHARS", 8), - patch.object(ProxyLogging, "_fire_deferred_stream_logging"), + patch.object(ProxyLogging, "fire_deferred_stream_logging"), ): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): @@ -7804,7 +7804,7 @@ async def test_async_data_generator_checks_raw_sse_buffer_limit_after_complete_f with ( patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), patch("litellm.proxy.proxy_server._MAX_RAW_SSE_BUFFER_CHARS", 8), - patch.object(ProxyLogging, "_fire_deferred_stream_logging"), + patch.object(ProxyLogging, "fire_deferred_stream_logging"), ): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): @@ -7854,7 +7854,7 @@ async def test_async_data_generator_google_genai_stream_omits_openai_done(): mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -7896,7 +7896,7 @@ async def test_async_data_generator_does_not_mark_completed_stream_as_disconnect mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator( mock_response, @@ -7946,7 +7946,7 @@ async def test_async_data_generator_google_genai_stream_forwards_error_without_d mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(ProxyLogging, "fire_deferred_stream_logging"): yielded_data = [] async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data): yielded_data.append(data) @@ -9384,7 +9384,7 @@ async def test_window_spend_counter_skips_invalid_window_start(): @pytest.mark.asyncio async def test_window_spend_counter_does_not_seed_zero_when_db_unavailable(): from litellm.caching.dual_cache import DualCache - from litellm.proxy.proxy_server import _ensure_window_spend_counter_initialized + from litellm.proxy.proxy_server import ensure_window_spend_counter_initialized counter_cache = DualCache() counter_key = "spend:key:key-window-db-unavailable:window:1h" @@ -9395,7 +9395,7 @@ async def test_window_spend_counter_does_not_seed_zero_when_db_unavailable(): ps.spend_counter_cache = counter_cache ps.prisma_client = None try: - initialized = await _ensure_window_spend_counter_initialized( + initialized = await ensure_window_spend_counter_initialized( counter_key=counter_key, entity_type="Key", entity_id="key-window-db-unavailable", @@ -9563,7 +9563,7 @@ async def test_increment_spend_counters_reseeds_from_db_on_bad_reserved_counter( @pytest.mark.asyncio async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure(): from litellm.caching.dual_cache import DualCache - from litellm.proxy.proxy_server import _increment_spend_counter_cache + from litellm.proxy.proxy_server import increment_spend_counter_cache counter_cache = DualCache() counter_cache.in_memory_cache.set_cache(key="spend:team:redis-fail", value=4.0) @@ -9578,7 +9578,7 @@ async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure( ps.spend_counter_cache = counter_cache try: with pytest.raises(RuntimeError): - await _increment_spend_counter_cache( + await increment_spend_counter_cache( counter_key="spend:team:redis-fail", increment=0.5, ) @@ -11161,14 +11161,14 @@ async def _lit6463_drive_realtime_session_holding_a_max_parallel_slot( endpoint did to the slot.""" from litellm.proxy import proxy_server as ps from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, _request_stash, ) from litellm.proxy.utils import InternalUsageCache dual_cache: Final = DualCache() await dual_cache.async_set_cache(key=_LIT6463_COUNTER_KEY, value={"slot-1": 1.0, "slot-2": 2.0}, local_only=True) - limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache)) + limiter: Final = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache)) stash: Final = RequestRateLimiterStash(parallel_slot={"slot_id": "slot-1", "counter_keys": [_LIT6463_COUNTER_KEY]}) reservation: Final = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []} @@ -11271,7 +11271,7 @@ async def test_release_or_invalidate_falls_back_to_invalidating_the_counters(): br, "release_budget_reservation", new=AsyncMock(side_effect=RuntimeError("counter store down")) ) # test-quality-ok: forces the failure branch; assertion observes which counter key got invalidated sink = patch.object( - ps, "_invalidate_spend_counter", new=_record + ps, "invalidate_spend_counter", new=_record ) # test-quality-ok: fakes the counter-store sink so the invalidated key is observable with failing_release, sink: await br.release_or_invalidate_budget_reservation(budget_reservation=reservation) @@ -11518,7 +11518,7 @@ class TestDeleteDeploymentSync: mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=Exception("DB connection lost")) - result = await proxy_config._get_models_from_db(prisma_client=mock_prisma) + result = await proxy_config.get_models_from_db(prisma_client=mock_prisma) assert result is None, f"Expected None on DB failure to signal fetch error, got {result!r}" @@ -11548,7 +11548,7 @@ class TestDeleteDeploymentSync: reader=PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False), ) - result = await ProxyConfig()._get_models_from_db(prisma_client=mock_prisma) + result = await ProxyConfig().get_models_from_db(prisma_client=mock_prisma) assert result == [committed_row], f"Expected the writer's just-committed row, got {result!r}" reader_inner.litellm_proxymodeltable.find_many.assert_not_awaited() @@ -11587,7 +11587,7 @@ class TestDeleteDeploymentSync: ) mock_prisma.db._writer_unavailable = True - result = await ProxyConfig()._get_models_from_db(prisma_client=mock_prisma) + result = await ProxyConfig().get_models_from_db(prisma_client=mock_prisma) assert result == [replica_row], f"Expected the replica's rows in degraded mode, got {result!r}" writer_inner.litellm_proxymodeltable.find_many.assert_not_awaited() @@ -13755,7 +13755,7 @@ async def _collect_async_data_generator_frames(request_data: dict) -> list: mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): - with patch.object(proxy_server_module.ProxyLogging, "_fire_deferred_stream_logging"): + with patch.object(proxy_server_module.ProxyLogging, "fire_deferred_stream_logging"): return [ frame.decode("utf-8") if isinstance(frame, bytes) else frame async for frame in async_data_generator(MockStream(), MagicMock(spec=UserAPIKeyAuth), request_data) @@ -13842,7 +13842,7 @@ def test_startup_warns_when_mock_testing_params_enabled(caplog): ) with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings={MOCK_TESTING_CONFIG_KEY: True}) + ProxyStartupEvent.warn_if_mock_testing_params_enabled(general_settings={MOCK_TESTING_CONFIG_KEY: True}) assert MOCK_TESTING_CONFIG_KEY in caplog.text for param_name in GATED_MOCK_PARAM_NAMES: @@ -13857,7 +13857,7 @@ def test_startup_is_silent_when_mock_testing_params_disabled(caplog): from litellm.proxy.route_llm_request import MOCK_TESTING_CONFIG_KEY with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings={}) + ProxyStartupEvent.warn_if_mock_testing_params_enabled(general_settings={}) assert MOCK_TESTING_CONFIG_KEY not in caplog.text @@ -14932,7 +14932,7 @@ class TestEmbeddingsFailureHookRequestData: with ( patch.object( proxy_server_module, - "_read_request_body", + "read_request_body", new=AsyncMock(return_value={"model": "my-embed", "input": "hello"}), ), patch.object( diff --git a/tests/unit/proxy/test_proxy_setting_guardrails.py b/tests/unit/proxy/test_proxy_setting_guardrails.py index c1d2c640b93..a734dbdbd7d 100644 --- a/tests/unit/proxy/test_proxy_setting_guardrails.py +++ b/tests/unit/proxy/test_proxy_setting_guardrails.py @@ -46,7 +46,7 @@ def test_active_callbacks(client): expected_callback_names = [ "lakeraAI_Moderation", - "_OPTIONAL_PromptInjectionDetectio", + "OPTIONAL_PromptInjectionDetection", "_ENTERPRISE_SecretDetection", ] diff --git a/tests/unit/proxy/test_proxy_token_counter.py b/tests/unit/proxy/test_proxy_token_counter.py index fcd72e3777c..e7a32816464 100644 --- a/tests/unit/proxy/test_proxy_token_counter.py +++ b/tests/unit/proxy/test_proxy_token_counter.py @@ -167,8 +167,8 @@ async def test_anthropic_messages_count_tokens_endpoint(): # Patch the _read_request_body function import litellm.proxy.anthropic_endpoints.endpoints as anthropic_endpoints - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body # Mock the internal token_counter function to return a controlled response async def mock_token_counter(request, call_endpoint=False): @@ -207,7 +207,7 @@ async def test_anthropic_messages_count_tokens_endpoint(): finally: # Restore original functions - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter @@ -241,8 +241,8 @@ async def test_anthropic_messages_count_tokens_with_non_anthropic_model(): # Patch the _read_request_body function import litellm.proxy.anthropic_endpoints.endpoints as anthropic_endpoints - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body # Mock the internal token_counter function to return a controlled response async def mock_token_counter(request, call_endpoint=True): @@ -281,7 +281,7 @@ async def test_anthropic_messages_count_tokens_with_non_anthropic_model(): finally: # Restore original functions - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter @@ -381,8 +381,8 @@ async def test_anthropic_endpoint_error_handling(): import litellm.proxy.anthropic_endpoints.endpoints as anthropic_endpoints - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body try: # Should raise HTTPException for missing model @@ -395,7 +395,7 @@ async def test_anthropic_endpoint_error_handling(): print("✅ Error handling test passed!") finally: - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body @pytest.mark.asyncio @@ -1111,8 +1111,8 @@ async def test_anthropic_endpoint_returns_anthropic_error_format(): mock_user_api_key_dict = MagicMock() - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body original_token_counter = proxy_server.token_counter @@ -1140,7 +1140,7 @@ async def test_anthropic_endpoint_returns_anthropic_error_format(): assert detail["error"]["type"] == "invalid_request_error" assert detail["error"]["message"] == "Input is too long for requested model." finally: - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter @@ -1163,8 +1163,8 @@ async def test_anthropic_endpoint_403_permission_error_format(): mock_user_api_key_dict = MagicMock() - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body original_token_counter = proxy_server.token_counter @@ -1190,7 +1190,7 @@ async def test_anthropic_endpoint_403_permission_error_format(): assert detail["error"]["type"] == "permission_error" assert detail["error"]["message"] == "Bearer Token has expired" finally: - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter @@ -1213,8 +1213,8 @@ async def test_anthropic_endpoint_429_rate_limit_error_format(): mock_user_api_key_dict = MagicMock() - original_read_request_body = anthropic_endpoints._read_request_body - anthropic_endpoints._read_request_body = mock_read_request_body + original_read_request_body = anthropic_endpoints.read_request_body + anthropic_endpoints.read_request_body = mock_read_request_body original_token_counter = proxy_server.token_counter @@ -1240,5 +1240,5 @@ async def test_anthropic_endpoint_429_rate_limit_error_format(): assert detail["error"]["type"] == "rate_limit_error" assert detail["error"]["message"] == "Rate limit exceeded" finally: - anthropic_endpoints._read_request_body = original_read_request_body + anthropic_endpoints.read_request_body = original_read_request_body proxy_server.token_counter = original_token_counter diff --git a/tests/unit/proxy/test_proxy_utils.py b/tests/unit/proxy/test_proxy_utils.py index f52a5bd431a..ba61e393281 100644 --- a/tests/unit/proxy/test_proxy_utils.py +++ b/tests/unit/proxy/test_proxy_utils.py @@ -10,7 +10,7 @@ from fastapi import HTTPException, Request from starlette.datastructures import State from litellm.integrations.custom_guardrail import CustomGuardrail -from litellm.proxy.utils import _get_docs_url, _get_openapi_url, _get_redoc_url +from litellm.proxy.utils import get_docs_url, get_openapi_url, get_redoc_url from litellm.types.guardrails import GuardrailEventHooks from unittest.mock import AsyncMock, MagicMock, patch @@ -22,7 +22,7 @@ from litellm.proxy.auth.auth_utils import ( is_request_body_safe, ) from litellm.proxy.litellm_pre_call_utils import ( - _get_dynamic_logging_metadata, + get_dynamic_logging_metadata, add_litellm_data_to_request, ) from pydantic import ValidationError @@ -294,7 +294,7 @@ def test_dynamic_logging_metadata_key_and_team_metadata(callback_vars): rpm_limit_per_model=None, tpm_limit_per_model=None, ) - callbacks = _get_dynamic_logging_metadata( + callbacks = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) @@ -332,7 +332,7 @@ def test_dynamic_logging_metadata_ignores_env_references_from_key_metadata( team_metadata={}, ) - callbacks = _get_dynamic_logging_metadata( + callbacks = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) @@ -410,7 +410,7 @@ def test_dynamic_turn_off_message_logging(callback_vars): rpm_limit_per_model=None, tpm_limit_per_model=None, ) - callbacks = _get_dynamic_logging_metadata( + callbacks = get_dynamic_logging_metadata( user_api_key_dict=user_api_key_dict, proxy_config=proxy_config ) @@ -774,7 +774,7 @@ def test_get_redoc_url(env_vars, expected_url): for key, value in env_vars.items(): os.environ[key] = value - result = _get_redoc_url() + result = get_redoc_url() assert result == expected_url @@ -799,7 +799,7 @@ def test_get_docs_url(env_vars, expected_url): for key, value in env_vars.items(): os.environ[key] = value - result = _get_docs_url() + result = get_docs_url() assert result == expected_url @@ -824,7 +824,7 @@ def test_get_openapi_url(env_vars, expected_url): for key, value in env_vars.items(): os.environ[key] = value - result = _get_openapi_url() + result = get_openapi_url() assert result == expected_url @@ -1569,7 +1569,7 @@ def test_is_allowed_to_make_key_request(): def test_get_model_group_info(): from litellm import Router - from litellm.proxy.proxy_server import _get_model_group_info + from litellm.proxy.proxy_server import get_model_group_info router = Router( model_list=[ @@ -1589,7 +1589,7 @@ def test_get_model_group_info(): }, ] ) - model_list = _get_model_group_info( + model_list = get_model_group_info( llm_router=router, all_models_str=["openai/tts-1", "openai/gpt-3.5-turbo"], model_group="openai/tts-1", diff --git a/tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py b/tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py index 34260908a5c..ec7104fa459 100644 --- a/tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py +++ b/tests/unit/proxy/test_proxy_utils_model_creation_and_error_logging.py @@ -195,7 +195,7 @@ async def test_proxy_only_error_log_keeps_the_request_litellm_call_id(monkeypatc def test_get_model_group_info_order(): from litellm import Router - from litellm.proxy.proxy_server import _get_model_group_info + from litellm.proxy.proxy_server import get_model_group_info router = Router( model_list=[ @@ -215,7 +215,7 @@ def test_get_model_group_info_order(): }, ] ) - model_list = _get_model_group_info( + model_list = get_model_group_info( llm_router=router, all_models_str=["openai/tts-1", "openai/gpt-3.5-turbo"], model_group=None, @@ -277,10 +277,10 @@ def _patch_today(monkeypatch, year, month, day): def test_get_projected_spend_over_limit_day_one(monkeypatch): - from litellm.proxy.utils import _get_projected_spend_over_limit + from litellm.proxy.utils import get_projected_spend_over_limit _patch_today(monkeypatch, 2026, 1, 1) - result = _get_projected_spend_over_limit(100.0, 1.0) + result = get_projected_spend_over_limit(100.0, 1.0) assert result is not None projected_spend, projected_exceeded_date = result @@ -289,10 +289,10 @@ def test_get_projected_spend_over_limit_day_one(monkeypatch): def test_get_projected_spend_over_limit_december(monkeypatch): - from litellm.proxy.utils import _get_projected_spend_over_limit + from litellm.proxy.utils import get_projected_spend_over_limit _patch_today(monkeypatch, 2026, 12, 15) - result = _get_projected_spend_over_limit(100.0, 1.0) + result = get_projected_spend_over_limit(100.0, 1.0) assert result is not None projected_spend, projected_exceeded_date = result @@ -301,10 +301,10 @@ def test_get_projected_spend_over_limit_december(monkeypatch): def test_get_projected_spend_over_limit_includes_current_spend(monkeypatch): - from litellm.proxy.utils import _get_projected_spend_over_limit + from litellm.proxy.utils import get_projected_spend_over_limit _patch_today(monkeypatch, 2026, 4, 11) - result = _get_projected_spend_over_limit(100.0, 200.0) + result = get_projected_spend_over_limit(100.0, 200.0) assert result is not None projected_spend, projected_exceeded_date = result @@ -629,7 +629,7 @@ class TestPostCallFailureHookLiftsCallTypeAndStartTime: from unittest.mock import AsyncMock from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + from litellm.proxy.hooks.proxy_track_cost_callback import ProxyDBLogger from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload request_start = real_datetime.datetime.now() - real_datetime.timedelta(seconds=2) @@ -659,7 +659,7 @@ class TestPostCallFailureHookLiftsCallTypeAndStartTime: proxy_logging_obj.alert_types = [] spend_writer = SimpleNamespace(update_database=AsyncMock()) original_callbacks = list(litellm.callbacks) - litellm.callbacks = [_ProxyDBLogger(spend_writer=lambda: spend_writer)] + litellm.callbacks = [ProxyDBLogger(spend_writer=lambda: spend_writer)] try: with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()): await proxy_logging_obj.post_call_failure_hook( diff --git a/tests/unit/proxy/test_response_polling_pre_call_checks.py b/tests/unit/proxy/test_response_polling_pre_call_checks.py index 1dca00fd5fb..c6e4301a501 100644 --- a/tests/unit/proxy/test_response_polling_pre_call_checks.py +++ b/tests/unit/proxy/test_response_polling_pre_call_checks.py @@ -35,9 +35,7 @@ class TestSkipPreCallLogic: mock_proxy_logging.during_call_hook = AsyncMock() with ( - patch.object( - processor, "common_processing_pre_call_logic", new_callable=AsyncMock - ) as mock_pre_call, + patch.object(processor, "common_processing_pre_call_logic", new_callable=AsyncMock) as mock_pre_call, patch( "litellm.proxy.common_request_processing.route_request", new_callable=AsyncMock, @@ -119,7 +117,7 @@ class TestPollingEndpointPreCallGuard: generate_polling_id_mock = MagicMock(return_value="litellm_poll_test") proxy_server_patches = { - "litellm.proxy.proxy_server._read_request_body": AsyncMock( + "litellm.proxy.proxy_server.read_request_body": AsyncMock( return_value={"model": "gpt-4", "background": True} ), "litellm.proxy.proxy_server.general_settings": {}, @@ -155,15 +153,11 @@ class TestPollingEndpointPreCallGuard: ), patch.object( ProxyBaseLLMRequestProcessing, - "_handle_llm_api_exception", + "handle_llm_api_exception", new_callable=AsyncMock, - return_value=HTTPException( - status_code=429, detail="Rate limit exceeded" - ), - ), - patch.object( - ResponsePollingHandler, "generate_polling_id", generate_polling_id_mock + return_value=HTTPException(status_code=429, detail="Rate limit exceeded"), ), + patch.object(ResponsePollingHandler, "generate_polling_id", generate_polling_id_mock), # Prevent background task from running (avoids noise from incomplete mocks) patch("asyncio.create_task"), patch.object( diff --git a/tests/unit/proxy/test_team_member_update.py b/tests/unit/proxy/test_team_member_update.py index ace4c4e65af..d77904fc05d 100644 --- a/tests/unit/proxy/test_team_member_update.py +++ b/tests/unit/proxy/test_team_member_update.py @@ -91,21 +91,17 @@ def happy_path_upsert(monkeypatch): AsyncMock( return_value={ "team_info": team_row, - "team_memberships": [ - types.SimpleNamespace(user_id="user-1", budget_id="bud-1") - ], + "team_memberships": [types.SimpleNamespace(user_id="user-1", budget_id="bud-1")], } ), ) upsert_mock = AsyncMock() - monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock) + monkeypatch.setattr(team_endpoints, "upsert_budget_and_membership", upsert_mock) return upsert_mock def _member_update_request(**overrides): - data = TeamMemberUpdateRequest( - team_id="team-1234", user_id="user-1", role="user", **overrides - ) + data = TeamMemberUpdateRequest(team_id="team-1234", user_id="user-1", role="user", **overrides) request = Request({"type": "http", "method": "POST", "path": "/team/member_update"}) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin") return data, request, auth @@ -115,9 +111,7 @@ def _member_update_request(**overrides): async def test_team_member_update_sends_provided_fields_as_patch(happy_path_upsert): """Fields the request sets must reach _upsert_budget_and_membership as a budget patch, otherwise the member budget is never written/reset.""" - data, request, auth = _member_update_request( - max_budget_in_team=10.0, budget_duration="30d" - ) + data, request, auth = _member_update_request(max_budget_in_team=10.0, budget_duration="30d") response = await team_member_update(data, request, auth) @@ -137,9 +131,7 @@ async def test_team_member_update_explicit_null_clears_field(happy_path_upsert): await team_member_update(data, request, auth) - assert happy_path_upsert.await_args.kwargs["budget_patch"] == { - "budget_duration": None - } + assert happy_path_upsert.await_args.kwargs["budget_patch"] == {"budget_duration": None} @pytest.mark.asyncio @@ -163,15 +155,13 @@ async def test_team_member_update_omits_unset_fields_from_patch(happy_path_upser ], ) @pytest.mark.asyncio -async def test_team_member_update_rejects_invalid_budget_duration( - monkeypatch, bad_duration -): +async def test_team_member_update_rejects_invalid_budget_duration(monkeypatch, bad_duration): """An invalid budget_duration must be rejected with a 400 before any DB write, so it can never be persisted and later break the budget reset job.""" monkeypatch.setattr(proxy_server, "prisma_client", object()) monkeypatch.setattr(proxy_server, "premium_user", False) upsert_mock = AsyncMock() - monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock) + monkeypatch.setattr(team_endpoints, "upsert_budget_and_membership", upsert_mock) data = TeamMemberUpdateRequest( team_id="team-1234", diff --git a/tests/unit/proxy/test_unit_test_proxy_hooks.py b/tests/unit/proxy/test_unit_test_proxy_hooks.py index e6ffea35e52..6ce7154f2d1 100644 --- a/tests/unit/proxy/test_unit_test_proxy_hooks.py +++ b/tests/unit/proxy/test_unit_test_proxy_hooks.py @@ -2,7 +2,7 @@ import asyncio from unittest.mock import Mock, patch, AsyncMock import pytest from fastapi import Request -from litellm.proxy.utils import _get_redoc_url, _get_docs_url +from litellm.proxy.utils import get_redoc_url, get_docs_url from datetime import datetime import litellm diff --git a/tests/unit/proxy/test_update_spend.py b/tests/unit/proxy/test_update_spend.py index 6b92320762b..62adb7303e2 100644 --- a/tests/unit/proxy/test_update_spend.py +++ b/tests/unit/proxy/test_update_spend.py @@ -1,6 +1,6 @@ import asyncio from unittest.mock import Mock -from litellm.proxy.utils import _get_redoc_url, _get_docs_url +from litellm.proxy.utils import get_redoc_url, get_docs_url import pytest from fastapi import Request diff --git a/tests/unit/proxy/test_zero_cost_model_budget_bypass.py b/tests/unit/proxy/test_zero_cost_model_budget_bypass.py index 56133f2d35b..c2931470482 100644 --- a/tests/unit/proxy/test_zero_cost_model_budget_bypass.py +++ b/tests/unit/proxy/test_zero_cost_model_budget_bypass.py @@ -23,8 +23,8 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.auth_checks import ( _check_team_member_budget, - _is_model_cost_zero, - _team_max_budget_check, + is_model_cost_zero, + team_max_budget_check, common_checks, ) from litellm.proxy.utils import ProxyLogging @@ -103,7 +103,7 @@ class TestIsModelCostZero: def test_zero_cost_model_in_router(self, mock_router_with_zero_cost_model): """Test that a zero-cost model in router is correctly identified.""" - result = _is_model_cost_zero( + result = is_model_cost_zero( model="on-prem-model", llm_router=mock_router_with_zero_cost_model ) assert result is True @@ -116,26 +116,26 @@ class TestIsModelCostZero: "input_cost_per_token": 0.0000015, "output_cost_per_token": 0.000002, } - result = _is_model_cost_zero( + result = is_model_cost_zero( model="cloud-model", llm_router=mock_router_with_zero_cost_model ) assert result is False def test_none_model(self, mock_router_with_zero_cost_model): """Test that None model returns False.""" - result = _is_model_cost_zero( + result = is_model_cost_zero( model=None, llm_router=mock_router_with_zero_cost_model ) assert result is False def test_none_router(self): """Test that None router returns False.""" - result = _is_model_cost_zero(model="some-model", llm_router=None) + result = is_model_cost_zero(model="some-model", llm_router=None) assert result is False def test_list_of_zero_cost_models(self, mock_router_with_zero_cost_model): """Test that a list of zero-cost models returns True.""" - result = _is_model_cost_zero( + result = is_model_cost_zero( model=["on-prem-model"], llm_router=mock_router_with_zero_cost_model ) assert result is True @@ -147,7 +147,7 @@ class TestIsModelCostZero: "input_cost_per_token": 0.0000015, "output_cost_per_token": 0.000002, } - result = _is_model_cost_zero( + result = is_model_cost_zero( model=["on-prem-model", "cloud-model"], llm_router=mock_router_with_zero_cost_model, ) @@ -514,7 +514,7 @@ class TestEdgeCases: with patch("litellm.get_model_info") as mock_get_model_info: # Simulate model not found mock_get_model_info.side_effect = Exception("Model not found") - result = _is_model_cost_zero( + result = is_model_cost_zero( model="nonexistent-model", llm_router=mock_router_with_zero_cost_model ) # Should return False (conservative approach) diff --git a/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 4a7cee80c61..0cdf264d8fc 100644 --- a/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -445,7 +445,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_decrypt_and_set_db_env_variables", + "decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables, ) @@ -655,7 +655,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -739,7 +739,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -806,7 +806,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -874,7 +874,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -953,7 +953,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -1024,7 +1024,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -1092,7 +1092,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -1964,7 +1964,7 @@ class TestProxySettingEndpoints: from litellm.proxy.proxy_server import proxy_config - monkeypatch.setattr(proxy_config, "_encrypt_env_variables", mock_encrypt) + monkeypatch.setattr(proxy_config, "encrypt_env_variables", mock_encrypt) # New SSO settings to save new_sso_settings = { @@ -2043,7 +2043,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -2092,7 +2092,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -2127,7 +2127,7 @@ class TestProxySettingEndpoints: from litellm.proxy.proxy_server import proxy_config monkeypatch.setattr( - proxy_config, "_decrypt_and_set_db_env_variables", mock_decrypt_and_set + proxy_config, "decrypt_and_set_db_env_variables", mock_decrypt_and_set ) response = client.get("/get/sso_settings") @@ -2212,7 +2212,7 @@ class TestProxySettingEndpoints: return environment_variables monkeypatch.setattr( - proxy_config, "_decrypt_and_set_db_env_variables", mock_decrypt + proxy_config, "decrypt_and_set_db_env_variables", mock_decrypt ) response = client.get("/get/sso_settings") @@ -2257,7 +2257,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -2309,7 +2309,7 @@ class TestProxySettingEndpoints: ) monkeypatch.setattr( proxy_config, - "_decrypt_and_set_db_env_variables", + "decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables, ) @@ -2431,7 +2431,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr( proxy_config, - "_decrypt_and_set_db_env_variables", + "decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables, ) @@ -2581,7 +2581,7 @@ def test_update_sso_settings_writes_redacted_audit_log(mock_proxy_config, monkey monkeypatch.setattr(litellm, "store_audit_logs", True) monkeypatch.setattr( proxy_server_module.proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) @@ -2653,14 +2653,14 @@ def test_update_sso_settings_audit_captures_redacted_before_snapshot( monkeypatch.setattr(litellm, "store_audit_logs", True) monkeypatch.setattr( proxy_server_module.proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) # Pretend the stored value is already plaintext for the test (production # decrypts via Fernet); the audit helper still has to redact it. monkeypatch.setattr( proxy_server_module.proxy_config, - "_decrypt_db_variables", + "decrypt_db_variables", lambda variables_dict: dict(variables_dict), ) @@ -2934,7 +2934,7 @@ def test_update_ui_theme_settings_writes_audit_log(mock_proxy_config, monkeypatc monkeypatch.setattr(litellm, "store_audit_logs", True) monkeypatch.setattr( proxy_server_module.proxy_config, - "_encrypt_env_variables", + "encrypt_env_variables", lambda environment_variables: environment_variables, ) diff --git a/tests/unit/proxy/utils/helpers/test_guardrail_merge.py b/tests/unit/proxy/utils/helpers/test_guardrail_merge.py index be00e36acae..e505517eeff 100644 --- a/tests/unit/proxy/utils/helpers/test_guardrail_merge.py +++ b/tests/unit/proxy/utils/helpers/test_guardrail_merge.py @@ -4,7 +4,7 @@ from unittest.mock import MagicMock import pytest from litellm.proxy.utils import ( - _check_and_merge_model_level_guardrails, + check_and_merge_model_level_guardrails, _merge_guardrails_with_existing, ) @@ -55,7 +55,7 @@ def test_check_and_merge_model_level_guardrails_happy_path_merges_lists(): "guardrails": ["user-policy"], }, } - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) snapshot = { "model": result["model"], "model_info_id": result["metadata"]["model_info"]["id"], @@ -70,7 +70,7 @@ def test_check_and_merge_model_level_guardrails_happy_path_merges_lists(): def test_check_and_merge_model_level_guardrails_returns_data_when_router_none(): data = {"metadata": {"model_info": {"id": "x"}}, "model": "m", "other": 1} - result = _check_and_merge_model_level_guardrails(data, None) + result = check_and_merge_model_level_guardrails(data, None) assert result is data assert normalize(result) == { "metadata": {"model_info": {"id": "x"}}, @@ -84,7 +84,7 @@ def test_check_and_merge_model_level_guardrails_returns_data_when_model_id_missi deployment (router returns None for both lookups), data is unchanged.""" router = _router_with_deployment(["pii"]) # by_alias=False by default data = {"metadata": {"model_info": {}}, "model": "m", "extra": "v"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) snapshot = { "is_same_object": result is data, "metadata": result["metadata"], @@ -108,7 +108,7 @@ def test_check_and_merge_model_level_guardrails_falls_back_to_model_alias_when_m model alias (#29652) so DB/UI-assigned guardrails still fire.""" router = _router_with_deployment(["pii"], by_alias=True) data = {"metadata": {"model_info": {}}, "model": "m", "extra": "v"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) # Merge happened via the alias fallback. assert "pii" in result["metadata"]["guardrails"] router.get_model_list.assert_called_once() @@ -121,7 +121,7 @@ def test_check_and_merge_model_level_guardrails_unions_guardrails_across_group_d The fix is to union the guardrails from all deployments in the group.""" router = _router_with_deployments([["pii"], ["secret-scan"], None]) data = {"metadata": {"model_info": {}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert sorted(result["metadata"]["guardrails"]) == ["pii", "secret-scan"] @@ -130,7 +130,7 @@ def test_check_and_merge_model_level_guardrails_dedups_guardrails_across_group_d entries in the merged guardrails list.""" router = _router_with_deployments([["pii"], ["pii", "secret-scan"]]) data = {"metadata": {"model_info": {}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert sorted(result["metadata"]["guardrails"]) == ["pii", "secret-scan"] @@ -139,7 +139,7 @@ def test_check_and_merge_model_level_guardrails_group_with_no_guardrails_returns the helper returns the data unchanged.""" router = _router_with_deployments([None, None, []]) data = {"metadata": {"model_info": {}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert result is data assert "guardrails" not in result["metadata"] @@ -163,7 +163,7 @@ def test_check_and_merge_model_level_guardrails_ignores_client_model_info_id_whe "model": "guarded-alias", "metadata": {"model_info": {"id": "spoofed-unguarded-deployment"}}, } - result = _check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) + result = check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) assert "alias-secret-scan" in result["metadata"]["guardrails"] # The model_id lookup must NOT have been used. router.get_deployment.assert_not_called() @@ -179,7 +179,7 @@ def test_check_and_merge_model_level_guardrails_trusts_client_model_info_id_by_d "model": "any", "metadata": {"model_info": {"id": "deployment-123"}}, } - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert "post-call-guardrail" in result["metadata"]["guardrails"] router.get_deployment.assert_called_once_with(model_id="deployment-123") @@ -193,7 +193,7 @@ def test_check_and_merge_model_level_guardrails_post_call_accepts_bare_string_gu router = MagicMock() router.get_deployment.return_value = deployment data = {"model": "any", "metadata": {"model_info": {"id": "deployment-x"}}} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert "scalar-guardrail" in result["metadata"]["guardrails"] @@ -203,7 +203,7 @@ def test_check_and_merge_model_level_guardrails_alias_union_accepts_bare_string_ router.get_deployment.return_value = None router.get_model_list.return_value = [{"litellm_params": {"guardrails": "scalar-alias-guardrail"}}] data = {"model": "alias-m", "metadata": {"model_info": {}}} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert "scalar-alias-guardrail" in result["metadata"]["guardrails"] @@ -222,7 +222,7 @@ def test_check_and_merge_model_level_guardrails_alias_fallback_passes_team_id(): "user_api_key_team_id": "team-abc", }, } - result = _check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) + result = check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) assert "team-guardrail" in result["metadata"]["guardrails"] router.get_model_list.assert_called_once_with(model_name="team-scoped-alias", team_id="team-abc") @@ -238,21 +238,21 @@ def test_check_and_merge_model_level_guardrails_alias_fallback_reads_team_id_fro "metadata": {"model_info": {}}, "litellm_metadata": {"user_api_key_team_id": "team-xyz"}, } - _check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) + check_and_merge_model_level_guardrails(data, router, trust_client_model_info=False) router.get_model_list.assert_called_once_with(model_name="alias-m", team_id="team-xyz") def test_check_and_merge_model_level_guardrails_returns_data_when_deployment_none(): router = _router_without_deployment() data = {"metadata": {"model_info": {"id": "x"}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert result is data def test_check_and_merge_model_level_guardrails_returns_data_when_guardrails_none(): router = _router_with_deployment(None) data = {"metadata": {"model_info": {"id": "x"}}, "model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) assert result is data @@ -260,7 +260,7 @@ def test_check_and_merge_model_level_guardrails_handles_missing_metadata(): """No metadata at all + alias unknown to the router → data unchanged.""" router = _router_with_deployment(["pii"]) # by_alias=False data = {"model": "m"} - result = _check_and_merge_model_level_guardrails(data, router) + result = check_and_merge_model_level_guardrails(data, router) snapshot = { "is_same_object": result is data, "model": result["model"], @@ -277,7 +277,7 @@ def test_check_and_merge_model_level_guardrails_raises_when_metadata_is_not_dict router = _router_with_deployment(["pii"]) data = {"metadata": "not-a-dict", "model": "m"} with pytest.raises(AttributeError): - _check_and_merge_model_level_guardrails(data, router) + check_and_merge_model_level_guardrails(data, router) def test_merge_guardrails_with_existing_happy_path_combines_lists(): diff --git a/tests/unit/proxy/utils/helpers/test_month_end_projection.py b/tests/unit/proxy/utils/helpers/test_month_end_projection.py index 5afe1f4faf8..8c1a4ae746d 100644 --- a/tests/unit/proxy/utils/helpers/test_month_end_projection.py +++ b/tests/unit/proxy/utils/helpers/test_month_end_projection.py @@ -4,8 +4,8 @@ import pytest from litellm.proxy.utils import ( _get_month_end_date, - _get_projected_spend_over_limit, - _is_projected_spend_over_limit, + get_projected_spend_over_limit, + is_projected_spend_over_limit, ) @@ -59,9 +59,7 @@ def test_get_month_end_date_raises_on_non_date_input(): def test_is_projected_spend_over_limit_happy_path_under_budget(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) summary = { - "result": _is_projected_spend_over_limit( - current_spend=10.0, soft_budget_limit=1_000_000.0 - ), + "result": is_projected_spend_over_limit(current_spend=10.0, soft_budget_limit=1_000_000.0), "current_spend": 10.0, "soft_budget_limit": 1_000_000.0, } @@ -75,9 +73,7 @@ def test_is_projected_spend_over_limit_happy_path_under_budget(monkeypatch): def test_is_projected_spend_over_limit_happy_path_over_budget(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) summary = { - "result": _is_projected_spend_over_limit( - current_spend=100.0, soft_budget_limit=50.0 - ), + "result": is_projected_spend_over_limit(current_spend=100.0, soft_budget_limit=50.0), "current_spend": 100.0, "soft_budget_limit": 50.0, } @@ -91,9 +87,7 @@ def test_is_projected_spend_over_limit_happy_path_over_budget(monkeypatch): def test_is_projected_spend_over_limit_first_of_month_no_division_by_zero(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 1)) summary = { - "result": _is_projected_spend_over_limit( - current_spend=5.0, soft_budget_limit=10.0 - ), + "result": is_projected_spend_over_limit(current_spend=5.0, soft_budget_limit=10.0), "current_spend": 5.0, "soft_budget_limit": 10.0, } @@ -105,10 +99,7 @@ def test_is_projected_spend_over_limit_first_of_month_no_division_by_zero(monkey def test_is_projected_spend_over_limit_none_limit_returns_false(): - assert ( - _is_projected_spend_over_limit(current_spend=10_000.0, soft_budget_limit=None) - is False - ) + assert is_projected_spend_over_limit(current_spend=10_000.0, soft_budget_limit=None) is False def test_is_projected_spend_over_limit_raises_when_today_missing(monkeypatch): @@ -119,14 +110,12 @@ def test_is_projected_spend_over_limit_raises_when_today_missing(monkeypatch): monkeypatch.setattr("litellm.proxy.utils.date", _Broken) with pytest.raises(RuntimeError): - _is_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) + is_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) def test_get_projected_spend_over_limit_happy_path_over_budget(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) - result = _get_projected_spend_over_limit( - current_spend=100.0, soft_budget_limit=50.0 - ) + result = get_projected_spend_over_limit(current_spend=100.0, soft_budget_limit=50.0) assert result is not None projected, exceed_date = result summary = { @@ -147,7 +136,7 @@ def test_get_projected_spend_over_limit_first_of_month_uses_current_as_daily( monkeypatch, ): _freeze_today(monkeypatch, date(2024, 1, 1)) - result = _get_projected_spend_over_limit(current_spend=5.0, soft_budget_limit=10.0) + result = get_projected_spend_over_limit(current_spend=5.0, soft_budget_limit=10.0) assert result is not None projected, exceed_date = result expected_exceed = date(2024, 1, 1) + timedelta(days=1.0) @@ -167,7 +156,7 @@ def test_get_projected_spend_over_limit_first_of_month_uses_current_as_daily( def test_get_projected_spend_over_limit_zero_daily_spend_exceed_today(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) - result = _get_projected_spend_over_limit(current_spend=0.0, soft_budget_limit=-1.0) + result = get_projected_spend_over_limit(current_spend=0.0, soft_budget_limit=-1.0) assert result is not None projected, exceed_date = result summary = { @@ -184,17 +173,12 @@ def test_get_projected_spend_over_limit_zero_daily_spend_exceed_today(monkeypatc def test_get_projected_spend_over_limit_under_budget_returns_none(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) - assert ( - _get_projected_spend_over_limit( - current_spend=1.0, soft_budget_limit=1_000_000.0 - ) - is None - ) + assert get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1_000_000.0) is None def test_get_projected_spend_over_limit_exceed_date_uses_remaining_budget(monkeypatch): _freeze_today(monkeypatch, date(2024, 1, 11)) - result = _get_projected_spend_over_limit(current_spend=20.0, soft_budget_limit=30.0) + result = get_projected_spend_over_limit(current_spend=20.0, soft_budget_limit=30.0) assert result is not None projected, exceed_date = result daily = 20.0 / 10 @@ -215,10 +199,7 @@ def test_get_projected_spend_over_limit_exceed_date_uses_remaining_budget(monkey def test_get_projected_spend_over_limit_none_limit_returns_none(): - assert ( - _get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=None) - is None - ) + assert get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=None) is None def test_get_projected_spend_over_limit_raises_when_today_missing(monkeypatch): @@ -229,4 +210,4 @@ def test_get_projected_spend_over_limit_raises_when_today_missing(monkeypatch): monkeypatch.setattr("litellm.proxy.utils.date", _Broken) with pytest.raises(RuntimeError): - _get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) + get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) diff --git a/tests/unit/proxy/utils/helpers/test_premium_user_check.py b/tests/unit/proxy/utils/helpers/test_premium_user_check.py index 0a9539c6dc3..f78d634157e 100644 --- a/tests/unit/proxy/utils/helpers/test_premium_user_check.py +++ b/tests/unit/proxy/utils/helpers/test_premium_user_check.py @@ -1,7 +1,7 @@ import pytest from fastapi import HTTPException -from litellm.proxy.utils import _premium_user_check +from litellm.proxy.utils import premium_user_check def normalize(value): @@ -13,7 +13,7 @@ def test_premium_user_check_happy_path_no_raise_when_premium(monkeypatch): monkeypatch.setattr(ps, "premium_user", True, raising=False) summary = { - "result": _premium_user_check(), + "result": premium_user_check(), "premium_user": True, "raised": False, } @@ -29,7 +29,7 @@ def test_premium_user_check_happy_path_with_feature_no_raise(monkeypatch): monkeypatch.setattr(ps, "premium_user", True, raising=False) summary = { - "result": _premium_user_check(feature="model-routing"), + "result": premium_user_check(feature="model-routing"), "premium_user": True, "feature": "model-routing", } @@ -45,7 +45,7 @@ def test_premium_user_check_raises_when_not_premium(monkeypatch): monkeypatch.setattr(ps, "premium_user", False, raising=False) with pytest.raises(HTTPException) as exc_info: - _premium_user_check() + premium_user_check() snapshot = { "status_code": exc_info.value.status_code, "is_dict_detail": isinstance(exc_info.value.detail, dict), @@ -63,7 +63,7 @@ def test_premium_user_check_raises_with_feature_message(monkeypatch): monkeypatch.setattr(ps, "premium_user", False, raising=False) with pytest.raises(HTTPException) as exc_info: - _premium_user_check(feature="custom-callbacks") + premium_user_check(feature="custom-callbacks") error_msg = exc_info.value.detail["error"] snapshot = { "status_code": exc_info.value.status_code, diff --git a/tests/unit/proxy/utils/helpers/test_team_configs.py b/tests/unit/proxy/utils/helpers/test_team_configs.py index 185d4d26ff4..2a9b36b92b8 100644 --- a/tests/unit/proxy/utils/helpers/test_team_configs.py +++ b/tests/unit/proxy/utils/helpers/test_team_configs.py @@ -1,6 +1,6 @@ import pytest -from litellm.proxy.utils import _is_valid_team_configs +from litellm.proxy.utils import is_valid_team_configs def normalize(value): @@ -11,7 +11,7 @@ def test_is_valid_team_configs_happy_path_allowed_model_mutates_config(): team_config = {"models": ["gpt-4o", "gpt-4o-mini"], "max_budget": 100.0} request_data = {"model": "gpt-4o"} snapshot = { - "result": _is_valid_team_configs( + "result": is_valid_team_configs( team_id="team-1", team_config=team_config, request_data=request_data, @@ -30,7 +30,7 @@ def test_is_valid_team_configs_no_models_key_is_noop(): team_config = {"max_budget": 100.0, "tpm_limit": 1000} request_data = {"model": "anything"} snapshot = { - "result": _is_valid_team_configs( + "result": is_valid_team_configs( team_id="team-1", team_config=team_config, request_data=request_data, @@ -48,7 +48,7 @@ def test_is_valid_team_configs_no_models_key_is_noop(): def test_is_valid_team_configs_short_circuits_when_team_id_none(): team_config = {"models": ["only-this"]} snapshot = { - "result": _is_valid_team_configs( + "result": is_valid_team_configs( team_id=None, team_config=team_config, request_data={"model": "anything-else"}, @@ -67,7 +67,7 @@ def test_is_valid_team_configs_raises_on_model_not_in_team_models(): team_config = {"models": ["gpt-4o"]} request_data = {"model": "claude-haiku"} with pytest.raises(Exception, match='claude-haiku\\. Valid models for team are') as exc_info: - _is_valid_team_configs( + is_valid_team_configs( team_id="team-1", team_config=team_config, request_data=request_data, diff --git a/tests/unit/proxy/utils/helpers/test_url_helpers.py b/tests/unit/proxy/utils/helpers/test_url_helpers.py index 5f23c5fc20e..ca7520b87dc 100644 --- a/tests/unit/proxy/utils/helpers/test_url_helpers.py +++ b/tests/unit/proxy/utils/helpers/test_url_helpers.py @@ -1,9 +1,9 @@ import pytest from litellm.proxy.utils import ( - _get_docs_url, - _get_openapi_url, - _get_redoc_url, + get_docs_url, + get_openapi_url, + get_redoc_url, get_custom_url, get_proxy_base_url, get_server_root_path, @@ -33,7 +33,7 @@ def _clear_url_env(monkeypatch): def test_get_redoc_url_default(monkeypatch): _clear_url_env(monkeypatch) summary = { - "result": _get_redoc_url(), + "result": get_redoc_url(), "redoc_url_env": None, "no_redoc_env": None, } @@ -48,7 +48,7 @@ def test_get_redoc_url_custom_env(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("REDOC_URL", "/custom-redoc") summary = { - "result": _get_redoc_url(), + "result": get_redoc_url(), "redoc_url_env": "/custom-redoc", "default_overridden": True, } @@ -62,13 +62,13 @@ def test_get_redoc_url_custom_env(monkeypatch): def test_get_redoc_url_disabled_returns_none_error_path(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("NO_REDOC", "True") - assert _get_redoc_url() is None + assert get_redoc_url() is None def test_get_docs_url_default(monkeypatch): _clear_url_env(monkeypatch) summary = { - "result": _get_docs_url(), + "result": get_docs_url(), "no_docs": None, "docs_url": None, } @@ -83,7 +83,7 @@ def test_get_docs_url_custom_env(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("DOCS_URL", "/api-docs") summary = { - "result": _get_docs_url(), + "result": get_docs_url(), "env": "/api-docs", "default_overridden": True, } @@ -97,13 +97,13 @@ def test_get_docs_url_custom_env(monkeypatch): def test_get_docs_url_disabled_returns_none_error_path(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("NO_DOCS", "True") - assert _get_docs_url() is None + assert get_docs_url() is None def test_get_openapi_url_default(monkeypatch): _clear_url_env(monkeypatch) summary = { - "result": _get_openapi_url(), + "result": get_openapi_url(), "no_openapi": None, "openapi_url": None, } @@ -118,7 +118,7 @@ def test_get_openapi_url_custom_env(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("OPENAPI_URL", "/api-schema") summary = { - "result": _get_openapi_url(), + "result": get_openapi_url(), "env": "/api-schema", "default_overridden": True, } @@ -132,7 +132,7 @@ def test_get_openapi_url_custom_env(monkeypatch): def test_get_openapi_url_disabled_returns_none_error_path(monkeypatch): _clear_url_env(monkeypatch) monkeypatch.setenv("NO_OPENAPI", "True") - assert _get_openapi_url() is None + assert get_openapi_url() is None @pytest.mark.parametrize( diff --git a/tests/unit/proxy/utils/prisma_and_spend/conftest.py b/tests/unit/proxy/utils/prisma_and_spend/conftest.py index e37a82a023b..aefb95516af 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/conftest.py +++ b/tests/unit/proxy/utils/prisma_and_spend/conftest.py @@ -359,7 +359,7 @@ def proxy_logging_with_redis(fake_redis: FakeRedisList) -> MagicMock: proxy_logging.db_spend_update_writer = MagicMock() proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock() buffer = RedisUpdateBuffer(redis_cache=fake_redis) - buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=True) + buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=True) proxy_logging.db_spend_update_writer.redis_update_buffer = buffer return proxy_logging diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_cache_user_row.py b/tests/unit/proxy/utils/prisma_and_spend/test_cache_user_row.py index d1270b60b19..a40e4422d49 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_cache_user_row.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_cache_user_row.py @@ -12,7 +12,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest -from litellm.proxy.utils import _cache_user_row +from litellm.proxy.utils import cache_user_row @pytest.mark.asyncio @@ -28,7 +28,7 @@ async def test_cache_user_row_caches_on_miss( db = MagicMock() db.get_data = AsyncMock(return_value=user_row) - result = await _cache_user_row("u1", mock_dual_cache, db) + result = await cache_user_row("u1", mock_dual_cache, db) cache_key = "u1_user_api_key_user_id" pinned = { "result": result, @@ -54,7 +54,7 @@ async def test_cache_user_row_skips_db_on_cache_hit( mock_dual_cache._store[cache_key] = "cached-blob" db = MagicMock() db.get_data = AsyncMock(return_value=None) - result = await _cache_user_row("u-hit", mock_dual_cache, db) + result = await cache_user_row("u-hit", mock_dual_cache, db) assert result is None assert db.get_data.await_count == 0 @@ -66,7 +66,7 @@ async def test_cache_user_row_skips_set_when_user_row_lacks_model_dump_json( user_row = SimpleNamespace(user_id="u2", spend=1.0) db = MagicMock() db.get_data = AsyncMock(return_value=user_row) - await _cache_user_row("u2", mock_dual_cache, db) + await cache_user_row("u2", mock_dual_cache, db) assert mock_dual_cache._store == {} assert mock_dual_cache.set_cache.call_count == 0 @@ -78,4 +78,4 @@ async def test_cache_user_row_propagates_db_error( db = MagicMock() db.get_data = AsyncMock(side_effect=RuntimeError("db down")) with pytest.raises(RuntimeError, match="db down"): - await _cache_user_row("u3", mock_dual_cache, db) + await cache_user_row("u3", mock_dual_cache, db) diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_password_helpers.py b/tests/unit/proxy/utils/prisma_and_spend/test_password_helpers.py index 3c028473479..f241cd00cf5 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_password_helpers.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_password_helpers.py @@ -21,7 +21,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from litellm.proxy.utils import ( - _hash_token_if_needed, + hash_token_if_needed, hash_password, hash_token, migrate_passwords_to_scrypt_async, @@ -116,9 +116,9 @@ def test_hash_token_if_needed_handles_sk_prefix() -> None: already_hashed = hashlib.sha256(plain.encode()).hexdigest() not_a_secret = "token-without-sk-prefix" actual = { - "sk_input_is_hashed": _hash_token_if_needed(plain) == already_hashed, - "non_sk_passthrough": _hash_token_if_needed(not_a_secret) == not_a_secret, - "double_hash_stable": _hash_token_if_needed(already_hashed) == already_hashed, + "sk_input_is_hashed": hash_token_if_needed(plain) == already_hashed, + "non_sk_passthrough": hash_token_if_needed(not_a_secret) == not_a_secret, + "double_hash_stable": hash_token_if_needed(already_hashed) == already_hashed, } assert actual == { "sk_input_is_hashed": True, @@ -129,7 +129,7 @@ def test_hash_token_if_needed_handles_sk_prefix() -> None: def test_hash_token_if_needed_error_on_non_string() -> None: with pytest.raises(AttributeError): - _hash_token_if_needed(None) # type: ignore[arg-type] + hash_token_if_needed(None) # pyright: ignore[reportArgumentType] # intentional invalid input checks the error # --------------------------------------------------------------------------- @@ -191,12 +191,10 @@ async def test_migrate_passwords_upgrades_only_plaintext_rows() -> None: result = await migrate_passwords_to_scrypt_async(pc) updated_user_ids = sorted( - call.kwargs["where"]["user_id"] - for call in pc.db.litellm_usertable.update.await_args_list + call.kwargs["where"]["user_id"] for call in pc.db.litellm_usertable.update.await_args_list ) new_password_prefixes = sorted( - call.kwargs["data"]["password"][:7] - for call in pc.db.litellm_usertable.update.await_args_list + call.kwargs["data"]["password"][:7] for call in pc.db.litellm_usertable.update.await_args_list ) outcome = { "message": result, @@ -216,8 +214,6 @@ async def test_migrate_passwords_upgrades_only_plaintext_rows() -> None: async def test_migrate_passwords_raises_on_db_failure() -> None: pc = MagicMock() pc.db = MagicMock() - pc.db.litellm_usertable.find_many = AsyncMock( - side_effect=RuntimeError("db unavailable") - ) + pc.db.litellm_usertable.find_many = AsyncMock(side_effect=RuntimeError("db unavailable")) with pytest.raises(RuntimeError, match="db unavailable"): await migrate_passwords_to_scrypt_async(pc) diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py index 7c4f4e582be..fc7ae531c34 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -23,8 +23,8 @@ import pytest from litellm.constants import REDIS_SPEND_LOGS_BUFFER_KEY from litellm.proxy.utils import ( MAX_SPEND_LOG_DRAIN_ITERATIONS, - _monitor_spend_logs_queue, - _raise_failed_update_spend_exception, + monitor_spend_logs_queue, + raise_failed_update_spend_exception, drain_spend_logs_queue, recover_parked_spend_logs, update_daily_tag_spend, @@ -111,15 +111,15 @@ async def test_update_daily_tag_spend_redis_path_when_buffered( writer = MagicMock() proxy_logging.db_spend_update_writer = writer writer.redis_update_buffer = MagicMock() - writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=True) - writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() - writer._commit_daily_tag_spend_to_db = AsyncMock() + writer.redis_update_buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=True) + writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() + writer.commit_daily_tag_spend_to_db = AsyncMock() await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging) - redis_kwargs = writer._commit_daily_tag_spend_to_db_with_redis.await_args.kwargs + redis_kwargs = writer.commit_daily_tag_spend_to_db_with_redis.await_args.kwargs pinned = { - "redis_calls": writer._commit_daily_tag_spend_to_db_with_redis.await_count, - "direct_calls": writer._commit_daily_tag_spend_to_db.await_count, + "redis_calls": writer.commit_daily_tag_spend_to_db_with_redis.await_count, + "direct_calls": writer.commit_daily_tag_spend_to_db.await_count, "redis_kwargs_keys": sorted(redis_kwargs.keys()), "redis_n_retries": redis_kwargs["n_retry_times"], } @@ -139,13 +139,13 @@ async def test_update_daily_tag_spend_direct_path_when_no_redis( writer = MagicMock() proxy_logging.db_spend_update_writer = writer writer.redis_update_buffer = MagicMock() - writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=False) - writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() - writer._commit_daily_tag_spend_to_db = AsyncMock() + writer.redis_update_buffer.should_commit_spend_updates_to_redis = MagicMock(return_value=False) + writer.commit_daily_tag_spend_to_db_with_redis = AsyncMock() + writer.commit_daily_tag_spend_to_db = AsyncMock() await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging) - assert writer._commit_daily_tag_spend_to_db.await_count == 1 - assert writer._commit_daily_tag_spend_to_db_with_redis.await_count == 0 + assert writer.commit_daily_tag_spend_to_db.await_count == 1 + assert writer.commit_daily_tag_spend_to_db_with_redis.await_count == 0 @pytest.mark.asyncio @@ -159,10 +159,10 @@ async def test_update_daily_tag_spend_logs_and_swallows_errors( proxy_logging = MagicMock() proxy_logging.db_spend_update_writer = MagicMock() proxy_logging.db_spend_update_writer.redis_update_buffer = MagicMock() - proxy_logging.db_spend_update_writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( + proxy_logging.db_spend_update_writer.redis_update_buffer.should_commit_spend_updates_to_redis = MagicMock( return_value=False ) - proxy_logging.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock( + proxy_logging.db_spend_update_writer.commit_daily_tag_spend_to_db = AsyncMock( side_effect=RuntimeError("commit boom") ) await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging) @@ -473,7 +473,7 @@ async def test_monitor_spend_logs_queue_invokes_job_when_queue_nonempty( monkeypatch.setattr(utils_mod, "update_spend_logs_job", _fake_job) with pytest.raises(asyncio.CancelledError): - await _monitor_spend_logs_queue( + await monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=proxy_logging, @@ -510,7 +510,7 @@ async def test_monitor_spend_logs_queue_swallows_errors_and_backs_off( mock_prisma_client._spend_log_transactions_lock = bad_lock with pytest.raises(asyncio.CancelledError): - await _monitor_spend_logs_queue( + await monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=proxy_logging, @@ -543,7 +543,7 @@ async def test_monitor_spend_logs_queue_flushes_as_soon_as_one_is_requested( monkeypatch.setattr(utils_mod, "update_spend_logs_job", _fake_job) monitor: Final = asyncio.create_task( - _monitor_spend_logs_queue( + monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=MagicMock(), @@ -589,7 +589,7 @@ def test_monitor_spend_logs_queue_flush_survives_an_earlier_event_loop( mock_prisma_client.spend_log_transactions = [] monitor: Final = asyncio.create_task( - _monitor_spend_logs_queue( + monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=MagicMock(), @@ -641,7 +641,7 @@ async def test_flush_requested_before_the_monitor_starts_costs_the_row_nothing( assert mock_prisma_client.spend_log_flush_requested is None monitor: Final = asyncio.create_task( - _monitor_spend_logs_queue( + monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=MagicMock(), @@ -661,7 +661,7 @@ def test_raise_failed_update_spend_exception_emits_failure_handler() -> None: async def _runner() -> Any: try: - _raise_failed_update_spend_exception( + raise_failed_update_spend_exception( e=RuntimeError("boom"), start_time=0.0, proxy_logging_obj=proxy_logging, @@ -701,7 +701,7 @@ def test_raise_failed_update_spend_exception_raises_original_error() -> None: proxy_logging.failure_handler = AsyncMock() async def _runner() -> None: - _raise_failed_update_spend_exception( + raise_failed_update_spend_exception( e=ValueError("specific"), start_time=0.0, proxy_logging_obj=proxy_logging, @@ -921,7 +921,7 @@ async def test_monitor_spend_logs_queue_pulls_parked_rows_before_each_flush( monkeypatch.setattr(utils_mod, "_wait_for_spend_log_flush_request", _poll) with pytest.raises(asyncio.CancelledError): - await _monitor_spend_logs_queue( + await monitor_spend_logs_queue( prisma_client=mock_prisma_client, db_writer_client=None, proxy_logging_obj=proxy_logging_with_redis, 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 2bf6a8d7e6c..b7fc8c6dbeb 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 @@ -454,8 +454,8 @@ 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 @@ -464,8 +464,8 @@ def test_every_pre_call_customlogger_is_deliberately_classified(): 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")), - ("detect_prompt_injection", _load("litellm.proxy.hooks.prompt_injection_detection", "_OPTIONAL_PromptInjectionDetection")), - ("azure_content_safety", _load("litellm.proxy.hooks.azure_content_safety", "_PROXY_AzureContentSafety")), + ("detect_prompt_injection", _load("litellm.proxy.hooks.prompt_injection_detection", "OPTIONAL_PromptInjectionDetection")), + ("azure_content_safety", _load("litellm.proxy.hooks.azure_content_safety", "PROXY_AzureContentSafety")), ): if cls is not None: registered[name] = cls diff --git a/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index 268000517d3..981fbfcf3e5 100644 --- a/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/unit/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -37,7 +37,7 @@ async def test_vector_store_search_forces_path_id_over_body_id(): request = _mock_request() with ( patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", new=AsyncMock( return_value={ "vector_store_id": "vs_body_victim", @@ -85,7 +85,7 @@ async def test_vector_store_file_create_forces_path_id_over_body_id(): request = _mock_request() with ( patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", new=AsyncMock( return_value={ "vector_store_id": "vs_body_victim", @@ -122,18 +122,14 @@ async def test_vector_store_file_list_resolves_managed_ids_and_cursors(): captured_data = {} provider_file_id: Final = "file-list-owned" - managed_file_data: Final = ( - SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( - "application/json", - "unified-file", - "managed-deployment", - provider_file_id, - "managed-deployment-id", - ) - ) - managed_file_id: Final = ( - base64.urlsafe_b64encode(managed_file_data.encode()).decode().rstrip("=") + managed_file_data: Final = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( + "application/json", + "unified-file", + "managed-deployment", + provider_file_id, + "managed-deployment-id", ) + managed_file_id: Final = base64.urlsafe_b64encode(managed_file_data.encode()).decode().rstrip("=") user_api_key_dict: Final = UserAPIKeyAuth(team_models=["team-openai"]) managed_file: Final[VectorStoreFileObject] = { "id": provider_file_id, @@ -171,9 +167,7 @@ async def test_vector_store_file_list_resolves_managed_ids_and_cursors(): "provider_resource_id,vs_provider_native;" "model_id,managed-deployment" ) - vector_store_id = ( - base64.urlsafe_b64encode(raw_vector_store_id.encode()).decode().rstrip("=") - ) + vector_store_id = base64.urlsafe_b64encode(raw_vector_store_id.encode()).decode().rstrip("=") request = _mock_request() request.method = "GET" @@ -221,9 +215,7 @@ async def test_vector_store_file_list_resolves_managed_ids_and_cursors(): assert captured_data["vector_store_id"] == "vs_provider_native" assert captured_data["api_key"] == "sk-managed-deployment" assert captured_data["model"] == "openai/managed-deployment" - llm_router.get_deployment_credentials_with_provider.assert_called_once_with( - model_id="managed-deployment" - ) + llm_router.get_deployment_credentials_with_provider.assert_called_once_with(model_id="managed-deployment") proxy_logging_obj.get_proxy_hook.assert_called_once_with("managed_files") resolver.assert_awaited_once_with( provider_file_ids=(provider_file_id,), @@ -247,7 +239,7 @@ async def test_vector_store_file_create_denies_other_team_path_store(): request = _mock_request() with ( patch( - "litellm.proxy.proxy_server._read_request_body", + "litellm.proxy.proxy_server.read_request_body", new=AsyncMock(return_value={"file_id": "file_123"}), ), patch.object(litellm, "vector_store_registry", mock_registry), @@ -282,7 +274,7 @@ async def test_rag_query_denies_nested_other_team_vector_store(): request = _mock_request() with ( patch( - "litellm.proxy.rag_endpoints.endpoints._read_request_body", + "litellm.proxy.rag_endpoints.endpoints.read_request_body", new=AsyncMock( return_value={ "model": "gpt-4o-mini", @@ -501,9 +493,7 @@ async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback new=cache_helper, ), ): - vector_store = await get_litellm_managed_vector_store( - vector_store_id="vs_cached" - ) + vector_store = await get_litellm_managed_vector_store(vector_store_id="vs_cached") assert vector_store is not None assert vector_store["vector_store_id"] == "vs_cached" @@ -518,9 +508,7 @@ async def test_get_managed_vector_store_fails_closed_on_lookup_error(): ) mock_registry = MagicMock() - mock_registry.get_litellm_managed_vector_store_from_registry.side_effect = ( - RuntimeError("registry unavailable") - ) + mock_registry.get_litellm_managed_vector_store_from_registry.side_effect = RuntimeError("registry unavailable") with patch.object(litellm, "vector_store_registry", mock_registry): with pytest.raises(HTTPException) as exc_info: @@ -624,8 +612,8 @@ async def test_azure_passthrough_denies_other_team_vector_store_index(): index_object.litellm_params.vector_store_name = "tenant-b-store" mock_index_registry = MagicMock() - mock_index_registry.is_vector_store_index.side_effect = ( - lambda vector_store_index_name: vector_store_index_name == "managed_index" + mock_index_registry.is_vector_store_index.side_effect = lambda vector_store_index_name: ( + vector_store_index_name == "managed_index" ) mock_index_registry.get_vector_store_index_by_name.return_value = index_object diff --git a/tests/unit/proxy/video_endpoints/test_endpoints.py b/tests/unit/proxy/video_endpoints/test_endpoints.py index c5996f95f54..5d84068fede 100644 --- a/tests/unit/proxy/video_endpoints/test_endpoints.py +++ b/tests/unit/proxy/video_endpoints/test_endpoints.py @@ -138,11 +138,11 @@ def harness(): stack.enter_context( patch.object( ProxyBaseLLMRequestProcessing, - "_handle_llm_api_exception", + "handle_llm_api_exception", handle_exc, ) ) - stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context(patch.object(endpoints, "read_request_body", read_body)) stack.enter_context( patch.object(endpoints, "batch_to_bytesio", batch_to_bytesio) ) 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 d215ce292aa..f43e15cb342 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -51,7 +51,7 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=_DummyMCPResult()), # Newer logging path calls this to enrich spend logs metadata - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -315,7 +315,7 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n ) from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm - _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=fake_server) + _msm.global_mcp_server_manager.get_mcp_server_from_tool_name = MagicMock(return_value=fake_server) tool_name = "my_deepwiki-read_wiki_structure" tool_calls = [ @@ -357,7 +357,7 @@ async def test_execute_tool_calls_reverse_maps_display_name(monkeypatch): ) from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm - _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=colliding_server) + _msm.global_mcp_server_manager.get_mcp_server_from_tool_name = MagicMock(return_value=colliding_server) _msm.global_mcp_server_manager.get_mcp_server_by_name = MagicMock(return_value=fake_server) tool_name = "browse_repo_docs" @@ -510,7 +510,7 @@ async def test_execute_tool_calls_applies_post_call_hook_content(monkeypatch): catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=result), - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -552,7 +552,7 @@ async def test_execute_tool_calls_returns_proxy_result_without_logging(monkeypat catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=result), - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -586,7 +586,7 @@ async def test_execute_tool_calls_passes_logging_details_to_proxy_hook(monkeypat catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=result), - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -622,7 +622,7 @@ async def test_execute_tool_calls_continues_when_post_call_logging_fails(monkeyp catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=result), - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( @@ -1370,7 +1370,7 @@ async def test_bridge_listing_leaves_the_callers_catalog_unchanged( with ( patch.dict(manager.tool_name_to_mcp_server_name_mapping), patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), - patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())), + patch.object(manager, "create_mcp_client", AsyncMock(return_value=object())), patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), ): try: @@ -1444,7 +1444,7 @@ async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch ] client: Final = AsyncMock() client.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")]) - manager._create_mcp_client = AsyncMock(return_value=client) + manager.create_mcp_client = AsyncMock(return_value=client) manager._fetch_tools_with_timeout = AsyncMock(return_value=upstream) guardrail: Final = _BridgeMetadataGuardrail() logger: Final = ProxyLogging(user_api_key_cache=DualCache()) diff --git a/tests/unit/responses/mcp/test_mcp_streaming_iterator.py b/tests/unit/responses/mcp/test_mcp_streaming_iterator.py index 3982081706f..52e2adbe34b 100644 --- a/tests/unit/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/unit/responses/mcp/test_mcp_streaming_iterator.py @@ -81,7 +81,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=call_tool, - _get_mcp_server_from_tool_name=MagicMock(return_value=None), + get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None), ) monkeypatch.setattr( diff --git a/tests/unit/test_private_usage_aliases.py b/tests/unit/test_private_usage_aliases.py index 412cae4eb0c..a5d74092d16 100644 --- a/tests/unit/test_private_usage_aliases.py +++ b/tests/unit/test_private_usage_aliases.py @@ -6,6 +6,1138 @@ from typing import Final, cast import pytest +PROXY_CLASS_NAME_ALIAS_CASES: Final = ( + ( + "litellm.proxy.hooks.dynamic_rate_limiter", + "_PROXY_DynamicRateLimitHandler", + "PROXY_DynamicRateLimitHandler", + ), + ( + "litellm.proxy.hooks.dynamic_rate_limiter_v3", + "_PROXY_DynamicRateLimitHandlerV3", + "PROXY_DynamicRateLimitHandlerV3", + ), + ( + "litellm.proxy.common_utils.config_sync_pubsub", + "_ConfigSyncPubSub", + "ConfigSyncPubSub", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.presidio", + "_OPTIONAL_PresidioPIIMasking", + "OPTIONAL_PresidioPIIMasking", + ), + ( + "litellm.proxy.hooks.prompt_injection_detection", + "_OPTIONAL_PromptInjectionDetection", + "OPTIONAL_PromptInjectionDetection", + ), + ( + "litellm.proxy.hooks.batch_redis_get", + "_PROXY_BatchRedisRequests", + "PROXY_BatchRedisRequests", + ), + ( + "litellm.proxy.hooks.azure_content_safety", + "_PROXY_AzureContentSafety", + "PROXY_AzureContentSafety", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.cisco_ai_defense.cisco_ai_defense_mcp", + "_CiscoAIDefenseMcpMixin", + "CiscoAIDefenseMcpMixin", + ), + ( + "litellm.proxy.hooks.cache_control_check", + "_PROXY_CacheControlCheck", + "PROXY_CacheControlCheck", + ), + ( + "litellm.proxy.hooks.max_budget_per_session_limiter", + "_PROXY_MaxBudgetPerSessionHandler", + "PROXY_MaxBudgetPerSessionHandler", + ), + ( + "litellm.proxy.hooks.max_iterations_limiter", + "_PROXY_MaxIterationsHandler", + "PROXY_MaxIterationsHandler", + ), + ( + "litellm.proxy.hooks.parallel_request_limiter", + "_PROXY_MaxParallelRequestsHandler", + "PROXY_MaxParallelRequestsHandler", + ), + ( + "litellm.proxy.hooks.parallel_request_limiter_v3", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.hooks.sensitive_data_routing", + "_PROXY_SensitiveDataRoutingHandler", + "PROXY_SensitiveDataRoutingHandler", + ), + ( + "litellm.proxy.hooks.batch_rate_limiter", + "_PROXY_BatchRateLimiter", + "PROXY_BatchRateLimiter", + ), + ( + "litellm.proxy.hooks.model_max_budget_limiter", + "_PROXY_VirtualKeyModelMaxBudgetLimiter", + "PROXY_VirtualKeyModelMaxBudgetLimiter", + ), + ( + "litellm.proxy.hooks.proxy_track_cost_callback", + "_ProxyDBLogger", + "ProxyDBLogger", + ), + ( + "enterprise.litellm_enterprise.proxy.hooks.managed_files", + "_PROXY_LiteLLMManagedFiles", + "PROXY_LiteLLMManagedFiles", + ), + ( + "enterprise.litellm_enterprise.proxy.hooks.managed_vector_stores", + "_PROXY_LiteLLMManagedVectorStores", + "PROXY_LiteLLMManagedVectorStores", + ), +) + +PACKAGE_EXPORT_ALIAS_CASES: Final = ( + ("litellm.proxy.hooks", "_PROXY_CacheControlCheck", "PROXY_CacheControlCheck"), + ( + "litellm.proxy.hooks", + "_PROXY_MaxBudgetPerSessionHandler", + "PROXY_MaxBudgetPerSessionHandler", + ), + ("litellm.proxy.hooks", "_PROXY_MaxIterationsHandler", "PROXY_MaxIterationsHandler"), + ("litellm.proxy.hooks", "_PROXY_MaxParallelRequestsHandler", "PROXY_MaxParallelRequestsHandler"), + ( + "litellm.proxy.hooks", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.hooks", + "_PROXY_SensitiveDataRoutingHandler", + "PROXY_SensitiveDataRoutingHandler", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_all_names_per_competitor", + "build_all_names_per_competitor", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_comparison_blocked_words", + "build_comparison_blocked_words", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_competitor_guardrail_definitions", + "build_competitor_guardrail_definitions", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_name_blocked_words", + "build_name_blocked_words", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_recommendation_blocked_words", + "build_recommendation_blocked_words", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_build_refinement_prompt", + "build_refinement_prompt", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_clean_competitor_line", + "clean_competitor_line", + ), + ( + "litellm.proxy.management_endpoints.policy_endpoints", + "_parse_variations_response", + "parse_variations_response", + ), +) + +MODULE_IMPORT_ALIAS_CASES: Final = ( + ( + "litellm.litellm_core_utils.custom_logger_registry", + "_PROXY_DynamicRateLimitHandler", + "PROXY_DynamicRateLimitHandler", + ), + ( + "litellm.litellm_core_utils.custom_logger_registry", + "_PROXY_DynamicRateLimitHandlerV3", + "PROXY_DynamicRateLimitHandlerV3", + ), + ( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp", + "_run_centralized_common_checks", + "run_centralized_common_checks", + ), + ( + "litellm.proxy._experimental.mcp_server.bridge_token_flow", + "_V2_GCM_PREFIX", + "V2_GCM_PREFIX", + ), + ( + "litellm.proxy._experimental.mcp_server.db", + "_get_salt_key", + "get_salt_key", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_bridge_mint_error_response", + "bridge_mint_error_response", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_extract_user_id_from_request", + "extract_user_id_from_request", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_finish_bridge_mint", + "finish_bridge_mint", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_prepare_bridge_mint", + "prepare_bridge_mint", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_prepare_bridge_refresh", + "prepare_bridge_refresh", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints", + "_reload_active_user_by_id", + "reload_active_user_by_id", + ), + ( + "litellm.proxy._experimental.mcp_server.mcp_server_manager", + "_is_mcp_admitted_user_subject", + "is_mcp_admitted_user_subject", + ), + ( + "litellm.proxy._experimental.mcp_server.mcp_server_manager", + "_redact_mcp_resource_url", + "redact_mcp_resource_url", + ), + ( + "litellm.proxy._experimental.mcp_server.oauth2_flow_backfill", + "_decode_oauth_payload", + "decode_oauth_payload", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_caller_authorization_fans_out", + "caller_authorization_fans_out", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_client_forwarded_authorization_headers", + "client_forwarded_authorization_headers", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_redact_mcp_resource_url", + "redact_mcp_resource_url", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_request_auth_header", + "request_auth_header", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_request_extra_headers", + "request_extra_headers", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_request_resolved_auth_headers", + "request_resolved_auth_headers", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_resolve_openapi_tool_auth", + "resolve_openapi_tool_auth", + ), + ( + "litellm.proxy._experimental.mcp_server.operations", + "_should_strip_caller_authorization", + "should_strip_caller_authorization", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES", + "UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_apply_toolset_scope", + "apply_toolset_scope", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_inherit_credentials_from_existing_server", + "inherit_credentials_from_existing_server", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_redact_mcp_resource_url", + "redact_mcp_resource_url", + ), + ( + "litellm.proxy._experimental.mcp_server.rest_endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy._experimental.mcp_server.server", + "_is_mcp_admitted_user_subject", + "is_mcp_admitted_user_subject", + ), + ( + "litellm.proxy._experimental.mcp_server.server", + "_mcp_active_toolset_id", + "mcp_active_toolset_id", + ), + ( + "litellm.proxy._experimental.mcp_server.server", + "_mcp_gateway_initialize_instructions", + "mcp_gateway_initialize_instructions", + ), + ( + "litellm.proxy._experimental.mcp_server.server", + "_mcp_gateway_server_name", + "mcp_gateway_server_name", + ), + ( + "litellm.proxy.anthropic_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.anthropic_endpoints.gateway_endpoints", + "_safe_set_request_parsed_body", + "safe_set_request_parsed_body", + ), + ( + "litellm.proxy.auth.auth_checks", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.auth.auth_checks", + "_safe_get_request_query_params", + "safe_get_request_query_params", + ), + ( + "litellm.proxy.auth.auth_exception_handler", + "_get_request_ip_address", + "get_request_ip_address", + ), + ( + "litellm.proxy.auth.fallback_budget", + "_is_model_cost_zero", + "is_model_cost_zero", + ), + ( + "litellm.proxy.auth.ip_address_utils", + "_get_request_ip_address", + "get_request_ip_address", + ), + ( + "litellm.proxy.auth.resolvers.store", + "_cache_key_object", + "cache_key_object", + ), + ( + "litellm.proxy.auth.resolvers.store", + "_copy_user_api_key_auth_for_cache", + "copy_user_api_key_auth_for_cache", + ), + ( + "litellm.proxy.auth.resolvers.store", + "_fetch_key_object_from_db_with_reconnect", + "fetch_key_object_from_db_with_reconnect", + ), + ( + "litellm.proxy.auth.route_checks", + "_user_is_org_admin", + "user_is_org_admin", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_cache_key_object", + "cache_key_object", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_can_object_call_model", + "can_object_call_model", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_check_end_user_budget", + "check_end_user_budget", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_delete_cache_key_object", + "delete_cache_key_object", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_get_user_role", + "get_user_role", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_is_model_cost_zero", + "is_model_cost_zero", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_is_user_proxy_admin", + "is_user_proxy_admin", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_realtime_request_body", + "realtime_request_body", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_safe_get_request_query_params", + "safe_get_request_query_params", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_safe_set_request_parsed_body", + "safe_set_request_parsed_body", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_team_member_max_budget_alert_check", + "team_member_max_budget_alert_check", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_virtual_key_max_budget_alert_check", + "virtual_key_max_budget_alert_check", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_virtual_key_max_budget_check", + "virtual_key_max_budget_check", + ), + ( + "litellm.proxy.auth.user_api_key_auth", + "_virtual_key_soft_budget_check", + "virtual_key_soft_budget_check", + ), + ( + "litellm.proxy.batches_endpoints.endpoints", + "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", + ), + ( + "litellm.proxy.batches_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.common_request_processing", + "_check_and_merge_model_level_guardrails", + "check_and_merge_model_level_guardrails", + ), + ( + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub", + "_ConfigSyncPubSub", + "ConfigSyncPubSub", + ), + ( + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub", + "_pubsub_capable_client", + "pubsub_capable_client", + ), + ( + "litellm.proxy.common_utils.key_rotation_manager", + "_calculate_key_rotation_time", + "calculate_key_rotation_time", + ), + ( + "litellm.proxy.common_utils.openai_endpoint_utils", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.container_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.custom_hooks.custom_ui_sso_hook", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.db.db_span", + "_is_exception_related_to_db", + "is_exception_related_to_db", + ), + ( + "litellm.proxy.fine_tuning_endpoints.endpoints", + "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", + ), + ( + "litellm.proxy.google_endpoints.agents_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.google_endpoints.agents_endpoints", + "_safe_get_request_query_params", + "safe_get_request_query_params", + ), + ( + "litellm.proxy.google_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.guardrails.guardrail_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.azure.prompt_shield", + "_RESPONSES_API_CALL_TYPES", + "RESPONSES_API_CALL_TYPES", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.azure.text_moderation", + "_RESPONSES_API_CALL_TYPES", + "RESPONSES_API_CALL_TYPES", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.cisco_ai_defense.cisco_ai_defense", + "_CiscoAIDefenseMcpMixin", + "CiscoAIDefenseMcpMixin", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.airline", + "_compile_marker", + "compile_marker", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.airline", + "_count_signals", + "count_signals", + ), + ( + "litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.airline", + "_word_boundary_match", + "word_boundary_match", + ), + ( + "litellm.proxy.guardrails.guardrail_registry", + "_OPTIONAL_PresidioPIIMasking", + "OPTIONAL_PresidioPIIMasking", + ), + ( + "litellm.proxy.health_endpoints._health_endpoints", + "_clean_endpoint_data", + "clean_endpoint_data", + ), + ( + "litellm.proxy.health_endpoints._health_endpoints", + "_update_litellm_params_for_health_check", + "update_litellm_params_for_health_check", + ), + ( + "litellm.proxy.hooks.dynamic_rate_limiter_v3", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.hooks.key_management_event_hooks", + "_hash_token_if_needed", + "hash_token_if_needed", + ), + ( + "litellm.proxy.hooks.proxy_track_cost_callback", + "_sanitize_error_information_for_spend_logs", + "sanitize_error_information_for_spend_logs", + ), + ( + "litellm.proxy.litellm_pre_call_utils", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.management_endpoints.access_group_endpoints", + "_cache_access_object", + "cache_access_object", + ), + ( + "litellm.proxy.management_endpoints.access_group_endpoints", + "_cache_key_object", + "cache_key_object", + ), + ( + "litellm.proxy.management_endpoints.access_group_endpoints", + "_cache_team_object", + "cache_team_object", + ), + ( + "litellm.proxy.management_endpoints.access_group_endpoints", + "_get_team_object_from_cache", + "get_team_object_from_cache", + ), + ( + "litellm.proxy.management_endpoints.auto_router_endpoints", + "_virtual_key_max_budget_check", + "virtual_key_max_budget_check", + ), + ( + "litellm.proxy.management_endpoints.budget_management_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.common_utils", + "_premium_user_check", + "premium_user_check", + ), + ( + "litellm.proxy.management_endpoints.credential_migration", + "_ALGO_AES_GCM", + "ALGO_AES_GCM", + ), + ( + "litellm.proxy.management_endpoints.credential_migration", + "_ENCRYPTION_ALGORITHM_SETTING", + "ENCRYPTION_ALGORITHM_SETTING", + ), + ( + "litellm.proxy.management_endpoints.credential_migration", + "_V2_GCM_PREFIX", + "V2_GCM_PREFIX", + ), + ( + "litellm.proxy.management_endpoints.credential_migration", + "_get_salt_key", + "get_salt_key", + ), + ( + "litellm.proxy.management_endpoints.customer_endpoints", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_endpoints.internal_user_endpoints", + "_check_permissions_caller_permission", + "check_permissions_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.internal_user_endpoints", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_endpoints.internal_user_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.jwt_key_mapping_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_add_model_to_db", + "add_model_to_db", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_check_disable_global_guardrails_caller_permission", + "check_disable_global_guardrails_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_check_passthrough_routes_caller_permission", + "check_passthrough_routes_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_delete_cache_key_object", + "delete_cache_key_object", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_hash_token_if_needed", + "hash_token_if_needed", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_is_master_key", + "is_master_key", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_set_object_metadata_field", + "set_object_metadata_field", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_team_member_has_permission", + "team_member_has_permission", + ), + ( + "litellm.proxy.management_endpoints.key_management_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.mcp_management_endpoints", + "_raise_if_not_oauth2", + "raise_if_not_oauth2", + ), + ( + "litellm.proxy.management_endpoints.mcp_management_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.management_endpoints.mcp_management_endpoints", + "_user_api_key_auth_builder", + "user_api_key_auth_builder", + ), + ( + "litellm.proxy.management_endpoints.mcp_management_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.model_management_endpoints", + "_refresh_cached_team", + "refresh_cached_team", + ), + ( + "litellm.proxy.management_endpoints.organization_endpoints", + "_set_object_metadata_field", + "set_object_metadata_field", + ), + ( + "litellm.proxy.management_endpoints.organization_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.prompt_cache_prediction", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.management_endpoints.prompt_cache_prediction", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.management_endpoints.scim.scim_v2", + "_delete_cache_key_object", + "delete_cache_key_object", + ), + ( + "litellm.proxy.management_endpoints.scim.scim_v2", + "_premium_user_check", + "premium_user_check", + ), + ( + "litellm.proxy.management_endpoints.scim.scim_v2", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.management_endpoints.session_endpoints", + "_persist_deleted_verification_tokens", + "persist_deleted_verification_tokens", + ), + ( + "litellm.proxy.management_endpoints.team_callback_endpoints", + "_CALLBACK_VAR_ENCRYPTED_PREFIX", + "CALLBACK_VAR_ENCRYPTED_PREFIX", + ), + ( + "litellm.proxy.management_endpoints.team_callback_endpoints", + "_get_validated_callback_metadata", + "get_validated_callback_metadata", + ), + ( + "litellm.proxy.management_endpoints.team_callback_endpoints", + "_refresh_cached_team", + "refresh_cached_team", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_cache_team_object", + "cache_team_object", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_check_disable_global_guardrails_caller_permission", + "check_disable_global_guardrails_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_check_passthrough_routes_caller_permission", + "check_passthrough_routes_caller_permission", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_set_object_metadata_field", + "set_object_metadata_field", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_team_member_has_permission", + "team_member_has_permission", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_update_metadata_fields", + "update_metadata_fields", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_upsert_budget_and_membership", + "upsert_budget_and_membership", + ), + ( + "litellm.proxy.management_endpoints.team_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.management_endpoints.ui_sso", + "_get_request_ip_address", + "get_request_ip_address", + ), + ( + "litellm.proxy.management_helpers.access_group_key_sync", + "_delete_cache_access_object", + "delete_cache_access_object", + ), + ( + "litellm.proxy.management_helpers.access_group_team_sync", + "_delete_cache_access_object", + "delete_cache_access_object", + ), + ( + "litellm.proxy.management_helpers.auto_router_permissions", + "_check_team_member_model_access", + "check_team_member_model_access", + ), + ( + "litellm.proxy.management_helpers.bulk_team_member_budgets", + "_upsert_budget_and_membership", + "upsert_budget_and_membership", + ), + ( + "litellm.proxy.management_helpers.bulk_user_creation", + "_check_permissions_caller_permission", + "check_permissions_caller_permission", + ), + ( + "litellm.proxy.management_helpers.bulk_user_creation", + "_set_object_permission", + "set_object_permission", + ), + ( + "litellm.proxy.management_helpers.bulk_user_deletion", + "_persist_deleted_verification_tokens", + "persist_deleted_verification_tokens", + ), + ( + "litellm.proxy.management_helpers.utils", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.openai_files_endpoints.files_endpoints", + "_is_base64_encoded_unified_file_id", + "is_base64_encoded_unified_file_id", + ), + ( + "litellm.proxy.openai_files_endpoints.files_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints", + "_get_bearer_token", + "get_bearer_token", + ), + ( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints", + "_safe_set_request_parsed_body", + "safe_set_request_parsed_body", + ), + ( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints", + "_get_dynamic_logging_metadata", + "get_dynamic_logging_metadata", + ), + ( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.proxy_server", + "_OPTIONAL_PromptInjectionDetection", + "OPTIONAL_PromptInjectionDetection", + ), + ( + "litellm.proxy.proxy_server", + "_PROXY_VirtualKeyModelMaxBudgetLimiter", + "PROXY_VirtualKeyModelMaxBudgetLimiter", + ), + ( + "litellm.proxy.proxy_server", + "_ProxyDBLogger", + "ProxyDBLogger", + ), + ( + "litellm.proxy.proxy_server", + "_add_model_to_db", + "add_model_to_db", + ), + ( + "litellm.proxy.proxy_server", + "_add_team_model_to_db", + "add_team_model_to_db", + ), + ( + "litellm.proxy.proxy_server", + "_cache_user_row", + "cache_user_row", + ), + ( + "litellm.proxy.proxy_server", + "_deduplicate_litellm_router_models", + "deduplicate_litellm_router_models", + ), + ( + "litellm.proxy.proxy_server", + "_fetch_global_spend_with_event_coordination", + "fetch_global_spend_with_event_coordination", + ), + ( + "litellm.proxy.proxy_server", + "_get_docs_url", + "get_docs_url", + ), + ( + "litellm.proxy.proxy_server", + "_get_openapi_url", + "get_openapi_url", + ), + ( + "litellm.proxy.proxy_server", + "_get_projected_spend_over_limit", + "get_projected_spend_over_limit", + ), + ( + "litellm.proxy.proxy_server", + "_get_redoc_url", + "get_redoc_url", + ), + ( + "litellm.proxy.proxy_server", + "_is_azure_model_router_request", + "is_azure_model_router_request", + ), + ( + "litellm.proxy.proxy_server", + "_is_projected_spend_over_limit", + "is_projected_spend_over_limit", + ), + ( + "litellm.proxy.proxy_server", + "_is_valid_team_configs", + "is_valid_team_configs", + ), + ( + "litellm.proxy.proxy_server", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.proxy_server", + "_realtime_request_body", + "realtime_request_body", + ), + ( + "litellm.proxy.proxy_server", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.proxy_server", + "_should_return_raw_model_name", + "should_return_raw_model_name", + ), + ( + "litellm.proxy.proxy_server", + "_user_has_admin_privileges", + "user_has_admin_privileges", + ), + ( + "litellm.proxy.proxy_server", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.rag_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.rag_endpoints.endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.realtime_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.response_api_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), + ( + "litellm.proxy.response_api_endpoints.endpoints", + "_safe_set_request_parsed_body", + "safe_set_request_parsed_body", + ), + ( + "litellm.proxy.search_endpoints.search_tool_registry", + "_get_salt_key", + "get_salt_key", + ), + ( + "litellm.proxy.spend_tracking.cloudzero_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.spend_tracking.vantage_endpoints", + "_user_has_admin_view", + "user_api_key_has_admin_view", + ), + ( + "litellm.proxy.utils", + "_PROXY_CacheControlCheck", + "PROXY_CacheControlCheck", + ), + ( + "litellm.proxy.utils", + "_PROXY_MaxParallelRequestsHandler", + "PROXY_MaxParallelRequestsHandler", + ), + ( + "litellm.proxy.utils", + "_PROXY_MaxParallelRequestsHandler_v3", + "PROXY_MaxParallelRequestsHandler_v3", + ), + ( + "litellm.proxy.utils", + "_PROXY_SensitiveDataRoutingHandler", + "PROXY_SensitiveDataRoutingHandler", + ), + ( + "litellm.proxy.utils", + "_is_exception_related_to_db", + "is_exception_related_to_db", + ), + ( + "litellm.proxy.vertex_ai_endpoints.langfuse_endpoints", + "_get_dynamic_logging_metadata", + "get_dynamic_logging_metadata", + ), + ( + "litellm.proxy.vertex_ai_endpoints.langfuse_endpoints", + "_safe_get_request_headers", + "safe_get_request_headers", + ), + ( + "litellm.proxy.video_endpoints.endpoints", + "_read_request_body", + "read_request_body", + ), +) + ALIAS_CASES: Final = ( ( "enterprise.enterprise_hooks.banned_keywords", @@ -1176,6 +2308,1843 @@ LLMS_ALIAS_CASES: Final = ( ("litellm.llms.watsonx.common_utils", "", "_generate_watsonx_token", "generate_watsonx_token", False), ("litellm.llms.watsonx.common_utils", "", "_get_api_params", "get_api_params", False), ) +PROXY_ALIAS_CASES: Final = ( + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + '', + '_is_mcp_admitted_user_subject', + 'is_mcp_admitted_user_subject', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_get_mcp_auth_header_from_headers', + 'get_mcp_auth_header_from_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_get_mcp_server_auth_headers_from_headers', + 'get_mcp_server_auth_headers_from_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_get_mcp_servers_from_access_groups', + 'get_mcp_servers_from_access_groups', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_get_oauth2_headers_from_headers', + 'get_oauth2_headers_from_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp', + 'MCPRequestHandler', + '_safe_get_headers_from_scope', + 'safe_get_headers_from_scope', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_bridge_mint_error_response', + 'bridge_mint_error_response', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_extract_user_id_from_request', + 'extract_user_id_from_request', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_finish_bridge_mint', + 'finish_bridge_mint', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_prepare_bridge_mint', + 'prepare_bridge_mint', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_prepare_bridge_refresh', + 'prepare_bridge_refresh', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.bridge_token_flow', + '', + '_reload_active_user_by_id', + 'reload_active_user_by_id', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.byok_oauth_endpoints', + '', + '_user_id_from_session_cookie', + 'user_id_from_session_cookie', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.db', + '', + '_decode_oauth_payload', + 'decode_oauth_payload', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.discoverable_endpoints', + '', + '_raise_if_not_oauth2', + 'raise_if_not_oauth2', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_context', + '', + '_mcp_active_toolset_id', + 'mcp_active_toolset_id', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_context', + '', + '_mcp_gateway_initialize_instructions', + 'mcp_gateway_initialize_instructions', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_context', + '', + '_mcp_gateway_server_name', + 'mcp_gateway_server_name', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES', + 'UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_caller_authorization_fans_out', + 'caller_authorization_fans_out', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_client_forwarded_authorization_headers', + 'client_forwarded_authorization_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_resolve_openapi_tool_auth', + 'resolve_openapi_tool_auth', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + '', + '_should_strip_caller_authorization', + 'should_strip_caller_authorization', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_build_mcp_server_table', + 'build_mcp_server_table', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_build_stdio_env', + 'build_stdio_env', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_create_mcp_client', + 'create_mcp_client', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_ensure_upstream_initialize_instructions_cached', + 'ensure_upstream_initialize_instructions_cached', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_extract_subject_token', + 'extract_subject_token', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_get_mcp_server_from_tool_name', + 'get_mcp_server_from_tool_name', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_get_tools_from_server', + 'get_tools_from_server', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.mcp_server_manager', + 'MCPServerManager', + '_is_server_accessible_from_ip', + 'is_server_accessible_from_ip', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.oauth2_token_cache', + '', + '_compute_per_user_token_ttl', + 'compute_per_user_token_ttl', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.oauth_utils', + '', + '_redact_mcp_resource_url', + 'redact_mcp_resource_url', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.server', + '', + '_redact_mcp_resource_url', + 'redact_mcp_resource_url', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator', + '', + '_OPENAPI_TOOL_NAME_MAX_LEN', + 'OPENAPI_TOOL_NAME_MAX_LEN', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator', + '', + '_request_auth_header', + 'request_auth_header', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator', + '', + '_request_extra_headers', + 'request_extra_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator', + '', + '_request_resolved_auth_headers', + 'request_resolved_auth_headers', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.semantic_tool_filter', + 'SemanticMCPToolFilter', + '_extract_tool_info', + 'extract_tool_info', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.server', + '', + '_apply_toolset_scope', + 'apply_toolset_scope', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.server_resolution', + 'MCPServerRegistry', + '_build_mcp_server_table', + 'build_mcp_server_table', + False, + ), + ( + 'litellm.proxy._experimental.mcp_server.server_resolution', + 'MCPServerRegistry', + '_is_server_accessible_from_ip', + 'is_server_accessible_from_ip', + False, + ), + ( + 'litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_marketplace', + '', + '_get_prisma_client', + 'get_prisma_client', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_cache_access_object', + 'cache_access_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_cache_key_object', + 'cache_key_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_cache_team_object', + 'cache_team_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_can_object_call_model', + 'can_object_call_model', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_check_end_user_budget', + 'check_end_user_budget', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_check_model_access_helper', + 'check_model_access_helper', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_check_team_member_model_access', + 'check_team_member_model_access', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_copy_user_api_key_auth_for_cache', + 'copy_user_api_key_auth_for_cache', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_delete_cache_access_object', + 'delete_cache_access_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_delete_cache_key_object', + 'delete_cache_key_object', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_fetch_key_object_from_db_with_reconnect', + 'fetch_key_object_from_db_with_reconnect', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_agent_ids_from_access_groups', + 'get_agent_ids_from_access_groups', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_mcp_server_ids_from_access_groups', + 'get_mcp_server_ids_from_access_groups', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_models_from_access_groups', + 'get_models_from_access_groups', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_team_object_from_cache', + 'get_team_object_from_cache', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_get_user_role', + 'get_user_role', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_is_model_cost_zero', + 'is_model_cost_zero', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_is_user_proxy_admin', + 'is_user_proxy_admin', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_key_access_group_grants_model', + 'key_access_group_grants_model', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_organization_max_budget_check', + 'organization_max_budget_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_team_max_budget_check', + 'team_max_budget_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_team_member_max_budget_alert_check', + 'team_member_max_budget_alert_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_virtual_key_max_budget_alert_check', + 'virtual_key_max_budget_alert_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_virtual_key_max_budget_check', + 'virtual_key_max_budget_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks', + '', + '_virtual_key_soft_budget_check', + 'virtual_key_soft_budget_check', + False, + ), + ( + 'litellm.proxy.auth.auth_checks_organization', + '', + '_user_is_org_admin', + 'user_is_org_admin', + False, + ), + ( + 'litellm.proxy.auth.auth_exception_handler', + 'UserAPIKeyAuthExceptionHandler', + '_handle_authentication_error', + 'handle_authentication_error', + False, + ), + ( + 'litellm.proxy.auth.auth_utils', + '', + '_get_request_ip_address', + 'get_request_ip_address', + False, + ), + ( + 'litellm.proxy.auth.resolvers.store', + 'IdentityStore', + '_principal_from_key', + 'principal_from_key', + False, + ), + ( + 'litellm.proxy.auth.route_checks', + 'RouteChecks', + '_get_request_method', + 'get_request_method', + False, + ), + ( + 'litellm.proxy.auth.route_checks', + 'RouteChecks', + '_is_assistants_api_request', + 'is_assistants_api_request', + False, + ), + ( + 'litellm.proxy.auth.route_checks', + 'RouteChecks', + '_is_wildcard_pattern', + 'is_wildcard_pattern', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_enforce_key_and_fallback_model_access', + 'enforce_key_and_fallback_model_access', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_fetch_global_spend_with_event_coordination', + 'fetch_global_spend_with_event_coordination', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_get_bearer_token', + 'get_bearer_token', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_run_centralized_common_checks', + 'run_centralized_common_checks', + False, + ), + ( + 'litellm.proxy.auth.user_api_key_auth', + '', + '_user_api_key_auth_builder', + 'user_api_key_auth_builder', + False, + ), + ( + 'litellm.proxy.common_request_processing', + '', + '_is_azure_model_router_request', + 'is_azure_model_router_request', + False, + ), + ( + 'litellm.proxy.common_request_processing', + '', + '_should_return_raw_model_name', + 'should_return_raw_model_name', + False, + ), + ( + 'litellm.proxy.common_request_processing', + 'ProxyBaseLLMRequestProcessing', + '_finalize_streaming_generator_cleanup', + 'finalize_streaming_generator_cleanup', + False, + ), + ( + 'litellm.proxy.common_request_processing', + 'ProxyBaseLLMRequestProcessing', + '_handle_llm_api_exception', + 'handle_llm_api_exception', + False, + ), + ( + 'litellm.proxy.common_request_processing', + 'ProxyBaseLLMRequestProcessing', + '_process_chunk_with_cost_injection', + 'process_chunk_with_cost_injection', + False, + ), + ( + 'litellm.proxy.common_utils.callback_utils', + '', + '_CALLBACK_VAR_ENCRYPTED_PREFIX', + 'CALLBACK_VAR_ENCRYPTED_PREFIX', + False, + ), + ( + 'litellm.proxy.common_utils.config_sync_pubsub', + '', + '_ConfigSyncPubSub', + 'ConfigSyncPubSub', + False, + ), + ( + 'litellm.proxy.common_utils.config_sync_pubsub', + '', + '_pubsub_capable_client', + 'pubsub_capable_client', + False, + ), + ( + 'litellm.proxy.common_utils.encrypt_decrypt_utils', + '', + '_ALGO_AES_GCM', + 'ALGO_AES_GCM', + False, + ), + ( + 'litellm.proxy.common_utils.encrypt_decrypt_utils', + '', + '_ENCRYPTION_ALGORITHM_SETTING', + 'ENCRYPTION_ALGORITHM_SETTING', + False, + ), + ( + 'litellm.proxy.common_utils.encrypt_decrypt_utils', + '', + '_V2_GCM_PREFIX', + 'V2_GCM_PREFIX', + False, + ), + ( + 'litellm.proxy.common_utils.encrypt_decrypt_utils', + '', + '_get_salt_key', + 'get_salt_key', + False, + ), + ( + 'litellm.proxy.common_utils.http_parsing_utils', + '', + '_read_request_body', + 'read_request_body', + False, + ), + ( + 'litellm.proxy.common_utils.http_parsing_utils', + '', + '_safe_get_request_headers', + 'safe_get_request_headers', + False, + ), + ( + 'litellm.proxy.common_utils.http_parsing_utils', + '', + '_safe_get_request_query_params', + 'safe_get_request_query_params', + False, + ), + ( + 'litellm.proxy.common_utils.http_parsing_utils', + '', + '_safe_set_request_parsed_body', + 'safe_set_request_parsed_body', + False, + ), + ( + 'litellm.proxy.common_utils.realtime_utils', + '', + '_realtime_request_body', + 'realtime_request_body', + False, + ), + ( + 'litellm.proxy.db.db_spend_update_writer', + 'DBSpendUpdateWriter', + '_commit_daily_tag_spend_to_db', + 'commit_daily_tag_spend_to_db', + False, + ), + ( + 'litellm.proxy.db.db_spend_update_writer', + 'DBSpendUpdateWriter', + '_commit_daily_tag_spend_to_db_with_redis', + 'commit_daily_tag_spend_to_db_with_redis', + False, + ), + ( + 'litellm.proxy.db.db_spend_update_writer', + 'DBSpendUpdateWriter', + '_handle_spend_update_failure', + 'handle_spend_update_failure', + False, + ), + ( + 'litellm.proxy.db.db_transaction_queue.redis_update_buffer', + 'RedisUpdateBuffer', + '_should_commit_spend_updates_to_redis', + 'should_commit_spend_updates_to_redis', + False, + ), + ( + 'litellm.proxy.db.log_db_metrics', + '', + '_is_exception_related_to_db', + 'is_exception_related_to_db', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.azure.base', + '', + '_RESPONSES_API_CALL_TYPES', + 'RESPONSES_API_CALL_TYPES', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.cisco_ai_defense.cisco_ai_defense_mcp', + '', + '_CiscoAIDefenseMcpMixin', + 'CiscoAIDefenseMcpMixin', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base', + '', + '_compile_marker', + 'compile_marker', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base', + '', + '_count_signals', + 'count_signals', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base', + '', + '_word_boundary_match', + 'word_boundary_match', + False, + ), + ( + 'litellm.proxy.guardrails.guardrail_hooks.presidio', + '', + '_OPTIONAL_PresidioPIIMasking', + 'OPTIONAL_PresidioPIIMasking', + False, + ), + ( + 'litellm.proxy.health_check', + '', + '_clean_endpoint_data', + 'clean_endpoint_data', + False, + ), + ( + 'litellm.proxy.health_check', + '', + '_update_litellm_params_for_health_check', + 'update_litellm_params_for_health_check', + False, + ), + ( + 'litellm.proxy.health_endpoints._health_endpoints', + '', + '_convert_health_check_to_dict', + 'convert_health_check_to_dict', + False, + ), + ( + 'litellm.proxy.health_endpoints._health_endpoints', + '', + '_save_background_health_checks_to_db', + 'save_background_health_checks_to_db', + False, + ), + ( + 'litellm.proxy.hooks.azure_content_safety', + '', + '_PROXY_AzureContentSafety', + 'PROXY_AzureContentSafety', + False, + ), + ( + 'litellm.proxy.hooks.batch_rate_limiter', + '', + '_PROXY_BatchRateLimiter', + 'PROXY_BatchRateLimiter', + False, + ), + ( + 'litellm.proxy.hooks.batch_redis_get', + '', + '_PROXY_BatchRedisRequests', + 'PROXY_BatchRedisRequests', + False, + ), + ( + 'litellm.proxy.hooks.cache_control_check', + '', + '_PROXY_CacheControlCheck', + 'PROXY_CacheControlCheck', + False, + ), + ( + 'litellm.proxy.hooks.dynamic_rate_limiter', + '', + '_PROXY_DynamicRateLimitHandler', + 'PROXY_DynamicRateLimitHandler', + False, + ), + ( + 'litellm.proxy.hooks.dynamic_rate_limiter_v3', + '', + '_PROXY_DynamicRateLimitHandlerV3', + 'PROXY_DynamicRateLimitHandlerV3', + False, + ), + ( + 'litellm.proxy.hooks.max_budget_per_session_limiter', + '', + '_PROXY_MaxBudgetPerSessionHandler', + 'PROXY_MaxBudgetPerSessionHandler', + False, + ), + ( + 'litellm.proxy.hooks.max_iterations_limiter', + '', + '_PROXY_MaxIterationsHandler', + 'PROXY_MaxIterationsHandler', + False, + ), + ( + 'litellm.proxy.hooks.model_max_budget_limiter', + '', + '_PROXY_VirtualKeyModelMaxBudgetLimiter', + 'PROXY_VirtualKeyModelMaxBudgetLimiter', + False, + ), + ( + 'litellm.proxy.hooks.parallel_request_limiter', + '', + '_PROXY_MaxParallelRequestsHandler', + 'PROXY_MaxParallelRequestsHandler', + False, + ), + ( + 'litellm.proxy.hooks.parallel_request_limiter_v3', + '', + '_PROXY_MaxParallelRequestsHandler_v3', + 'PROXY_MaxParallelRequestsHandler_v3', + False, + ), + ( + 'litellm.proxy.hooks.parallel_request_limiter_v3', + '_PROXY_MaxParallelRequestsHandler_v3', + '_create_rate_limit_descriptors', + 'create_rate_limit_descriptors', + False, + ), + ( + 'litellm.proxy.hooks.prompt_injection_detection', + '', + '_OPTIONAL_PromptInjectionDetection', + 'OPTIONAL_PromptInjectionDetection', + False, + ), + ( + 'litellm.proxy.hooks.proxy_track_cost_callback', + '', + '_ProxyDBLogger', + 'ProxyDBLogger', + False, + ), + ( + 'litellm.proxy.hooks.sensitive_data_routing', + '', + '_PROXY_SensitiveDataRoutingHandler', + 'PROXY_SensitiveDataRoutingHandler', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + '', + '_add_guardrails_from_key_or_team_metadata', + 'add_guardrails_from_key_or_team_metadata', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + '', + '_get_dynamic_logging_metadata', + 'get_dynamic_logging_metadata', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + '', + '_get_metadata_variable_name', + 'get_metadata_variable_name', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + '', + '_get_validated_callback_metadata', + 'get_validated_callback_metadata', + False, + ), + ( + 'litellm.proxy.litellm_pre_call_utils', + 'LiteLLMProxyRequestSetup', + '_merge_tags', + 'merge_tags', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_check_disable_global_guardrails_caller_permission', + 'check_disable_global_guardrails_caller_permission', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_check_passthrough_routes_caller_permission', + 'check_passthrough_routes_caller_permission', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_set_object_metadata_field', + 'set_object_metadata_field', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_team_member_has_permission', + 'team_member_has_permission', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_update_metadata_fields', + 'update_metadata_fields', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_upsert_budget_and_membership', + 'upsert_budget_and_membership', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_user_has_admin_privileges', + 'user_has_admin_privileges', + False, + ), + ( + 'litellm.proxy.management_endpoints.common_utils', + '', + '_user_has_admin_view', + 'user_api_key_has_admin_view', + False, + ), + ( + 'litellm.proxy.management_endpoints.config_override_endpoints', + '', + '_clear_hashicorp_vault_state', + 'clear_hashicorp_vault_state', + False, + ), + ( + 'litellm.proxy.management_endpoints.config_override_endpoints', + '', + '_get_current_env_values', + 'get_current_env_values', + False, + ), + ( + 'litellm.proxy.management_endpoints.config_override_endpoints', + '', + '_parse_config_value', + 'parse_config_value', + False, + ), + ( + 'litellm.proxy.management_endpoints.config_override_endpoints', + '', + '_set_env_vars', + 'set_env_vars', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_calculate_key_rotation_time', + 'calculate_key_rotation_time', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_check_permissions_caller_permission', + 'check_permissions_caller_permission', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_get_caller_team_role', + 'get_caller_team_role', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_list_key_helper', + 'list_key_helper', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_persist_deleted_verification_tokens', + 'persist_deleted_verification_tokens', + False, + ), + ( + 'litellm.proxy.management_endpoints.key_management_endpoints', + '', + '_rotate_master_key', + 'rotate_master_key', + False, + ), + ( + 'litellm.proxy.management_endpoints.mcp_management_endpoints', + '', + '_inherit_credentials_from_existing_server', + 'inherit_credentials_from_existing_server', + False, + ), + ( + 'litellm.proxy.management_endpoints.model_management_endpoints', + '', + '_add_model_to_db', + 'add_model_to_db', + False, + ), + ( + 'litellm.proxy.management_endpoints.model_management_endpoints', + '', + '_add_team_model_to_db', + 'add_team_model_to_db', + False, + ), + ( + 'litellm.proxy.management_endpoints.model_management_endpoints', + '', + '_deduplicate_litellm_router_models', + 'deduplicate_litellm_router_models', + False, + ), + ( + 'litellm.proxy.management_endpoints.organization_endpoints', + '', + '_verify_org_access', + 'verify_org_access', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_all_names_per_competitor', + 'build_all_names_per_competitor', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_comparison_blocked_words', + 'build_comparison_blocked_words', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_competitor_guardrail_definitions', + 'build_competitor_guardrail_definitions', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_name_blocked_words', + 'build_name_blocked_words', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_recommendation_blocked_words', + 'build_recommendation_blocked_words', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_build_refinement_prompt', + 'build_refinement_prompt', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_clean_competitor_line', + 'clean_competitor_line', + False, + ), + ( + 'litellm.proxy.management_endpoints.policy_endpoints.endpoints', + '', + '_parse_variations_response', + 'parse_variations_response', + False, + ), + ( + 'litellm.proxy.management_endpoints.team_endpoints', + '', + '_cleanup_members_with_roles', + 'cleanup_members_with_roles', + False, + ), + ( + 'litellm.proxy.management_endpoints.team_endpoints', + '', + '_refresh_cached_team', + 'refresh_cached_team', + False, + ), + ( + 'litellm.proxy.management_endpoints.team_endpoints', + 'TeamMemberBudgetHandler', + '_clean_team_member_fields', + 'clean_team_member_fields', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + '', + '_sso_return_to_redirect', + 'sso_return_to_redirect', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_delete_pkce_verifier', + 'delete_pkce_verifier', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_get_cli_state', + 'get_cli_state', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_get_user_email_and_id_from_result', + 'get_user_email_and_id_from_result', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_pkce_token_exchange', + 'pkce_token_exchange', + False, + ), + ( + 'litellm.proxy.management_endpoints.ui_sso', + 'SSOAuthenticationHandler', + '_validate_return_to', + 'validate_return_to', + False, + ), + ( + 'litellm.proxy.management_helpers.object_permission_utils', + '', + '_get_allow_all_keys_server_ids', + 'get_allow_all_keys_server_ids', + False, + ), + ( + 'litellm.proxy.management_helpers.object_permission_utils', + '', + '_get_team_allowed_mcp_servers', + 'get_team_allowed_mcp_servers', + False, + ), + ( + 'litellm.proxy.management_helpers.object_permission_utils', + '', + '_set_object_permission', + 'set_object_permission', + False, + ), + ( + 'litellm.proxy.openai_files_endpoints.common_utils', + '', + '_is_base64_encoded_unified_file_id', + 'is_base64_encoded_unified_file_id', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints', + '', + '_extract_model_from_bedrock_endpoint', + 'extract_model_from_bedrock_endpoint', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints', + 'BaseOpenAIPassThroughHandler', + '_base_openai_pass_through_handler', + 'base_openai_pass_through_handler', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler', + 'AnthropicPassthroughLoggingHandler', + '_build_complete_streaming_response', + 'build_complete_streaming_response', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler', + 'AnthropicPassthroughLoggingHandler', + '_build_usage_only_response_from_chunks', + 'build_usage_only_response_from_chunks', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler', + 'AnthropicPassthroughLoggingHandler', + '_handle_logging_anthropic_collected_chunks', + 'handle_logging_anthropic_collected_chunks', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.assembly_passthrough_logging_handler', + 'AssemblyAIPassthroughLoggingHandler', + '_get_assembly_base_url_from_region', + 'get_assembly_base_url_from_region', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.assembly_passthrough_logging_handler', + 'AssemblyAIPassthroughLoggingHandler', + '_get_assembly_region_from_url', + 'get_assembly_region_from_url', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.assembly_passthrough_logging_handler', + 'AssemblyAIPassthroughLoggingHandler', + '_should_log_request', + 'should_log_request', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler', + '', + '_is_openai_compatible_url', + 'is_openai_compatible_url', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler', + 'OpenAIPassthroughLoggingHandler', + '_handle_logging_openai_collected_chunks', + 'handle_logging_openai_collected_chunks', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler', + 'VertexPassthroughLoggingHandler', + '_handle_logging_vertex_collected_chunks', + 'handle_logging_vertex_collected_chunks', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.pass_through_endpoints', + 'HttpPassThroughEndpointHelpers', + '_init_kwargs_for_pass_through_endpoint', + 'init_kwargs_for_pass_through_endpoint', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.pass_through_endpoints', + 'HttpPassThroughEndpointHelpers', + '_update_stream_param_based_on_request_body', + 'update_stream_param_based_on_request_body', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.streaming_handler', + 'PassThroughStreamingHandler', + '_route_streaming_logging_to_handler', + 'route_streaming_logging_to_handler', + False, + ), + ( + 'litellm.proxy.pass_through_endpoints.success_handler', + 'PassThroughEndpointLogging', + '_handle_logging', + 'handle_logging', + False, + ), + ( + 'litellm.proxy.policy_engine.policy_registry', + 'PolicyRegistry', + '_parse_policy', + 'parse_policy', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_apply_uvicorn_max_requests_jitter', + 'apply_uvicorn_max_requests_jitter', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_configure_dev_reload', + 'configure_dev_reload', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_echo_litellm_version', + 'echo_litellm_version', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_get_default_unvicorn_init_args', + 'get_default_unvicorn_init_args', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_get_loop_type', + 'get_loop_type', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_init_granian_server', + 'init_granian_server', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_init_hypercorn_server', + 'init_hypercorn_server', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_is_port_in_use', + 'is_port_in_use', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_maybe_setup_prometheus_multiproc_dir', + 'maybe_setup_prometheus_multiproc_dir', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_config_validation', + 'run_config_validation', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_gunicorn_server', + 'run_gunicorn_server', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_health_check', + 'run_health_check', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_ollama_serve', + 'run_ollama_serve', + False, + ), + ( + 'litellm.proxy.proxy_cli', + 'ProxyInitializationHelpers', + '_run_test_chat_completion', + 'run_test_chat_completion', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_build_redis_usage_cache', + 'build_redis_usage_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_ensure_spend_counter_initialized', + 'ensure_spend_counter_initialized', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_ensure_window_spend_counter_initialized', + 'ensure_window_spend_counter_initialized', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_environment_has_redis_connection_target', + 'environment_has_redis_connection_target', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_get_model_group_info', + 'get_model_group_info', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_increment_spend_counter_cache', + 'increment_spend_counter_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_initialize_shared_aiohttp_session', + 'initialize_shared_aiohttp_session', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_invalidate_spend_counter', + 'invalidate_spend_counter', + False, + ), + ( + 'litellm.proxy.proxy_server', + '', + '_title', + 'title', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_add_deployment_locked', + 'add_deployment_locked', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_decrypt_and_set_db_env_variables', + 'decrypt_and_set_db_env_variables', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_decrypt_db_variables', + 'decrypt_db_variables', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_encrypt_env_variables', + 'encrypt_env_variables', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_encrypt_env_variables_for_db', + 'encrypt_env_variables_for_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_get_models_from_db', + 'get_models_from_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_init_cache', + 'init_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_init_semantic_filter_settings_in_db', + 'init_semantic_filter_settings_in_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyConfig', + '_serve_pass_through_endpoints', + 'serve_pass_through_endpoints', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_add_proxy_budget_to_db', + 'add_proxy_budget_to_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_attach_router_to_prompt_injection_detectors', + 'attach_router_to_prompt_injection_detectors', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_get_transaction_buffer_redis_cache', + 'get_transaction_buffer_redis_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_init_coordination_redis_from_db', + 'init_coordination_redis_from_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_init_dd_tracer', + 'init_dd_tracer', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_init_pyroscope', + 'init_pyroscope', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_initialize_jwt_auth', + 'initialize_jwt_auth', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_initialize_semantic_tool_filter', + 'initialize_semantic_tool_filter', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_initialize_startup_logging', + 'initialize_startup_logging', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_setup_prisma_client', + 'setup_prisma_client', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_sync_ui_settings_to_general_settings', + 'sync_ui_settings_to_general_settings', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_update_default_team_member_budget', + 'update_default_team_member_budget', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_validate_redis_transaction_buffer_config', + 'validate_redis_transaction_buffer_config', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_warm_global_spend_cache', + 'warm_global_spend_cache', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_warn_budget_without_db', + 'warn_budget_without_db', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_warn_fail_closed_rate_limits_without_redis', + 'warn_fail_closed_rate_limits_without_redis', + False, + ), + ( + 'litellm.proxy.proxy_server', + 'ProxyStartupEvent', + '_warn_if_mock_testing_params_enabled', + 'warn_if_mock_testing_params_enabled', + False, + ), + ( + 'litellm.proxy.spend_tracking.spend_management_endpoints', + '', + '_get_spend_report_for_time_range', + 'get_spend_report_for_time_range', + False, + ), + ( + 'litellm.proxy.spend_tracking.spend_management_endpoints', + '', + '_is_admin_view_safe', + 'is_admin_view_safe', + False, + ), + ( + 'litellm.proxy.spend_tracking.spend_tracking_utils', + '', + '_is_master_key', + 'is_master_key', + False, + ), + ( + 'litellm.proxy.spend_tracking.spend_tracking_utils', + '', + '_sanitize_error_information_for_spend_logs', + 'sanitize_error_information_for_spend_logs', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_cache_user_row', + 'cache_user_row', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_check_and_merge_model_level_guardrails', + 'check_and_merge_model_level_guardrails', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_get_docs_url', + 'get_docs_url', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_get_openapi_url', + 'get_openapi_url', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_get_projected_spend_over_limit', + 'get_projected_spend_over_limit', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_get_redoc_url', + 'get_redoc_url', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_hash_token_if_needed', + 'hash_token_if_needed', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_is_projected_spend_over_limit', + 'is_projected_spend_over_limit', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_is_valid_team_configs', + 'is_valid_team_configs', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_monitor_spend_logs_queue', + 'monitor_spend_logs_queue', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_premium_user_check', + 'premium_user_check', + False, + ), + ( + 'litellm.proxy.utils', + '', + '_raise_failed_update_spend_exception', + 'raise_failed_update_spend_exception', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_arelease_max_parallel_requests_on_disconnect', + 'arelease_max_parallel_requests_on_disconnect', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_callback_capabilities', + 'callback_capabilities', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_convert_mcp_hook_response_to_kwargs', + 'convert_mcp_hook_response_to_kwargs', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_create_mcp_request_object_from_kwargs', + 'create_mcp_request_object_from_kwargs', + False, + ), + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_fire_deferred_stream_logging', + 'fire_deferred_stream_logging', + False, + ), +) + LLMS_FORWARDER_CASES: Final = ( ( "litellm.llms.bedrock.chat.invoke_handler", @@ -2218,6 +5187,17 @@ LLMS_FORWARDER_CASES: Final = ( False, ), ) +PROXY_FORWARDER_CASES: Final = ( + ( + 'litellm.proxy.utils', + 'ProxyLogging', + '_convert_mcp_to_llm_format', + 'convert_mcp_to_llm_format', + 'instance', + False, + ), +) + def _get_owner(module_path: str, owner_name: str) -> object: @@ -2242,7 +5222,7 @@ def _get_instance(owner: object, public_name: str) -> object: @pytest.mark.parametrize( ("module_path", "owner_name", "old_name", "new_name", "use_instance"), - (*ALIAS_CASES, *LLMS_ALIAS_CASES), + (*ALIAS_CASES, *LLMS_ALIAS_CASES, *PROXY_ALIAS_CASES), ) def test_public_aliases( module_path: str, @@ -2260,6 +5240,51 @@ def test_public_aliases( assert old_value is new_value +@pytest.mark.parametrize( + ("package_name", "private_name", "public_name"), + PACKAGE_EXPORT_ALIAS_CASES, +) +def test_private_package_exports_are_available_and_match_public_alias( + package_name: str, + private_name: str, + public_name: str, +) -> None: + package: Final = import_module(package_name) + + assert getattr(package, private_name) is getattr(package, public_name) + + +@pytest.mark.parametrize( + ("module_path", "private_name", "public_name"), + MODULE_IMPORT_ALIAS_CASES, +) +def test_module_level_private_imports_remain_compatible( + module_path: str, + private_name: str, + public_name: str, +) -> None: + module: Final = import_module(module_path) + + assert getattr(module, private_name) is getattr(module, public_name) + + +@pytest.mark.parametrize( + ("module_path", "private_name", "public_name"), + PROXY_CLASS_NAME_ALIAS_CASES, +) +def test_proxy_class_aliases_keep_the_private_name( + module_path: str, + private_name: str, + public_name: str, +) -> None: + module: Final = import_module(module_path) + private_class: Final = cast(type[object], getattr(module, private_name)) + public_class: Final = cast(type[object], getattr(module, public_name)) + + assert public_class is private_class + assert public_class.__name__ == private_name + + def _make_private_override( descriptor: str, is_async: bool, @@ -2325,7 +5350,7 @@ def _forwarder_arguments( "descriptor", "is_async", ), - LLMS_FORWARDER_CASES, + (*LLMS_FORWARDER_CASES, *PROXY_FORWARDER_CASES), ) async def test_public_forwarders_dispatch_to_private_subclass_override( module_path: str, diff --git a/tests/unit/test_rate_limit_error_unification.py b/tests/unit/test_rate_limit_error_unification.py index e5acba938c7..ac2af7dba99 100644 --- a/tests/unit/test_rate_limit_error_unification.py +++ b/tests/unit/test_rate_limit_error_unification.py @@ -336,10 +336,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) with pytest.raises(ProxyRateLimitError) as exc_info: handler.raise_rate_limit_error(additional_details="key-over-rpm") e = exc_info.value @@ -366,10 +366,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) with pytest.raises(ProxyRateLimitError) as exc_info: handler.raise_rate_limit_error() # no additional_details detail_str = str(exc_info.value.detail) @@ -424,10 +424,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) - handler = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) # Minimal fabricated OVER_LIMIT response. The helper only reads a # handful of fields off `status` and ignores everything else. response = { @@ -475,13 +475,13 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.max_iterations_limiter import ( - _PROXY_MaxIterationsHandler, + PROXY_MaxIterationsHandler, ) from litellm.proxy.utils import InternalUsageCache from litellm.types.agents import AgentResponse cache = DualCache() - handler = _PROXY_MaxIterationsHandler( + handler = PROXY_MaxIterationsHandler( internal_usage_cache=InternalUsageCache(cache), ) user_api_key_dict = UserAPIKeyAuth( @@ -531,10 +531,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.dynamic_rate_limiter import ( - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) # check_available_usage returns (available_tpm, available_rpm, # model_tpm, model_rpm, active_projects). Setting available_tpm == 0 # forces the TPM-exceeded raise. @@ -570,12 +570,12 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) cache = MagicMock() cache.async_batch_set_cache = AsyncMock(return_value=None) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) with pytest.raises(ProxyRateLimitError) as exc_info: await handler.check_key_in_limits( user_api_key_dict=UserAPIKeyAuth(api_key="sk-key"), @@ -634,12 +634,12 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) cache = MagicMock() cache.async_batch_set_cache = AsyncMock(return_value=None) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) with pytest.raises(ProxyRateLimitError) as exc_info: await handler.check_key_in_limits( user_api_key_dict=UserAPIKeyAuth(api_key="sk-key"), @@ -692,12 +692,12 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) cache = MagicMock() cache.async_batch_set_cache = AsyncMock(return_value=None) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) with pytest.raises(ProxyRateLimitError) as exc_info: await handler.check_key_in_limits( user_api_key_dict=UserAPIKeyAuth(api_key="sk-key"), @@ -723,10 +723,10 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.dynamic_rate_limiter import ( - _PROXY_DynamicRateLimitHandler, + PROXY_DynamicRateLimitHandler, ) - handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) + handler = PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) # available_tpm > 0, available_rpm == 0 → RPM raise branch. handler.check_available_usage = AsyncMock( # type: ignore[method-assign] return_value=(100, 0, 1000, 100, 1) @@ -769,13 +769,13 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( - _PROXY_DynamicRateLimitHandlerV3, + PROXY_DynamicRateLimitHandlerV3, ) # Bypass __init__ — we want to inject a stub v3_limiter without # paying for the full handler setup. - handler = _PROXY_DynamicRateLimitHandlerV3.__new__( - _PROXY_DynamicRateLimitHandlerV3 + handler = PROXY_DynamicRateLimitHandlerV3.__new__( + PROXY_DynamicRateLimitHandlerV3 ) v3_limiter = MagicMock() v3_limiter.window_size = 60 @@ -836,12 +836,12 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.max_budget_per_session_limiter import ( - _PROXY_MaxBudgetPerSessionHandler, + PROXY_MaxBudgetPerSessionHandler, ) internal_cache = MagicMock() internal_cache.async_get_cache = AsyncMock(return_value=10.0) - handler = _PROXY_MaxBudgetPerSessionHandler( + handler = PROXY_MaxBudgetPerSessionHandler( internal_usage_cache=internal_cache, ) user_api_key_dict = UserAPIKeyAuth( @@ -876,14 +876,14 @@ class TestProxyHooksActuallyRaiseProxyRateLimitError: from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) # Inject a parallel_request_limiter mock with a usable window_size so # the helper's str(window_size) call doesn't NameError. parallel_limiter = MagicMock() parallel_limiter.window_size = 60 - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=parallel_limiter, ) @@ -1117,10 +1117,10 @@ class TestProxyHooksWireTypeCorrectly: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) with pytest.raises(ProxyRateLimitError) as exc_info: handler.raise_rate_limit_error() assert exc_info.value.rate_limit_type == "concurrent_requests" @@ -1129,10 +1129,10 @@ class TestProxyHooksWireTypeCorrectly: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter import ( - _PROXY_MaxParallelRequestsHandler, + PROXY_MaxParallelRequestsHandler, ) - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + handler = PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) with pytest.raises(ProxyRateLimitError) as exc_info: handler.raise_rate_limit_error( additional_details="tpm-zero", @@ -1174,10 +1174,10 @@ class TestProxyHooksWireTypeCorrectly: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) - handler = _PROXY_MaxParallelRequestsHandler_v3( + handler = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=MagicMock(), ) # Minimal RateLimitResponse + descriptors shape that the handler @@ -1225,10 +1225,10 @@ class TestProxyHooksWireTypeCorrectly: from unittest.mock import MagicMock from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - _PROXY_MaxParallelRequestsHandler_v3, + PROXY_MaxParallelRequestsHandler_v3, ) - handler = _PROXY_MaxParallelRequestsHandler_v3( + handler = PROXY_MaxParallelRequestsHandler_v3( internal_usage_cache=MagicMock(), ) response = { @@ -1269,12 +1269,12 @@ class TestProxyHooksWireTypeCorrectly: from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) prl = MagicMock() prl.window_size = 60 - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=prl, ) @@ -1312,12 +1312,12 @@ class TestProxyHooksWireTypeCorrectly: from litellm.proxy.hooks.batch_rate_limiter import ( BatchFileUsage, - _PROXY_BatchRateLimiter, + PROXY_BatchRateLimiter, ) prl = MagicMock() prl.window_size = 60 - handler = _PROXY_BatchRateLimiter( + handler = PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=prl, ) @@ -1558,7 +1558,7 @@ class TestBudgetExceededErrorLlmProviderEnrichment: {"use_x_forwarded_for": False}, ), patch( - "litellm.proxy.auth.auth_exception_handler._get_request_ip_address", + "litellm.proxy.auth.auth_exception_handler.get_request_ip_address", return_value="127.0.0.1", ), ): diff --git a/tests/unit/test_video_generation.py b/tests/unit/test_video_generation.py index dc15e82597d..5c4bc8ea652 100644 --- a/tests/unit/test_video_generation.py +++ b/tests/unit/test_video_generation.py @@ -2398,7 +2398,7 @@ async def test_edit_and_extension_read_cached_body_after_auth_consumes_stream( import litellm.proxy.video_endpoints.endpoints as endpoints from litellm.proxy._types import ProxyException, UserAPIKeyAuth - from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.common_utils.http_parsing_utils import read_request_body body = urlencode(form).encode() stream = {"sent": False} @@ -2423,7 +2423,7 @@ async def test_edit_and_extension_read_cached_body_after_auth_consumes_stream( receive, ) - await _read_request_body(request=request) + await read_request_body(request=request) handler = getattr(endpoints, handler_name) with pytest.raises(ProxyException) as exc_info: