From ef42461c1eaaf08ad14c1688700182602d96da02 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Mon, 26 May 2025 14:41:42 -0700 Subject: [PATCH] Litellm fix GitHub action testing (#11163) * test: add __init__.py files * refactor: rename test folder to avoid naming conflict * test: update workflows * test: update tests * test: update imports * test: update tests * test: remove unused import * ci(test-litellm.yml): add pytest retry to github workflow * test: fix test --- .github/workflows/test-litellm.yml | 3 +- .pre-commit-config.yaml | 2 +- ...odel_prices_and_context_window_backup.json | 4 +- .../anthropic_endpoints => }/__init__.py | 0 tests/litellm/integrations/test_athina.py | 207 ---------------- tests/litellm/proxy/client/cli/__init__.py | 1 - tests/test_litellm/__init__.py | 0 .../caching/test_in_memory_cache.py | 0 .../caching/test_redis_cache.py | 18 +- .../caching/test_redis_cluster_cache.py | 0 .../caching/test_redis_semantic_cache.py | 98 ++++---- tests/{litellm => test_litellm}/conftest.py | 0 .../send_emails/test_base_email.py | 0 .../send_emails/test_endpoints.py | 0 .../send_emails/test_resend_email.py | 0 .../experimental_mcp_client/test_tools.py | 0 .../SlackAlerting/test_slack_alerting.py | 6 +- .../test_slack_alerting_utils.py | 0 .../integrations/arize/test_arize_phoenix.py | 1 - .../integrations/arize/test_arize_utils.py | 11 +- .../integrations/gcs_pubsub/test_pub_sub.py | 0 .../integrations/test_agentops.py | 79 ++++--- .../test_anthropic_cache_control_hook.py | 0 .../test_litellm/integrations/test_athina.py | 220 ++++++++++++++++++ .../test_custom_prompt_management.py | 0 .../integrations/test_deepeval.py | 12 +- .../integrations/test_langfuse.py | 0 .../integrations/test_opentelemetry.py | 0 .../integrations/test_prometheus.py | 0 .../integrations/test_prometheus_services.py | 0 .../litellm_core_utils/__init__.py | 0 .../llm_cost_calc/test_llm_cost_calc_utils.py | 0 .../test_tool_call_cost_tracking.py | 0 .../messages_with_counts.py | 12 +- ...ore_utils_prompt_templates_common_utils.py | 0 ...llm_core_utils_prompt_templates_factory.py | 19 +- .../litellm_core_utils/test_core_helpers.py | 0 .../litellm_core_utils/test_dd_tracing.py | 0 .../test_duration_parser.py | 51 ++-- .../test_litellm_logging.py | 0 .../test_realtime_streaming.py | 0 .../test_safe_json_dumps.py | 0 .../test_streaming_chunk_builder_utils.py | 2 +- .../test_streaming_handler.py | 32 ++- .../litellm_core_utils/test_token_counter.py | 11 +- .../test_token_counter_tool.py | 7 +- .../test_token_counter_tool_data.py | 26 +-- .../chat/test_anthropic_chat_handler.py | 0 .../test_anthropic_chat_transformation.py | 0 .../test_azure_image_generation_init.py | 0 .../llms/azure/test_azure_common_utils.py | 0 .../chat/test_azure_ai_transformation.py | 16 +- .../chat/test_converse_transformation.py | 0 .../llms/bedrock/chat/test_invoke_handler.py | 0 .../llms/bedrock/chat/test_mistral_config.py | 10 +- .../test_amazon_stability3_transformation.py | 0 .../llms/bedrock/rerank/transformation.py | 0 .../llms/bedrock/test_base_aws_llm.py | 0 .../llms/bedrock/test_bedrock_common_utils.py | 0 .../llms/chat/test_converse_handler.py | 0 .../llms/cohere/chat/test_transformation.py | 12 +- .../llms/custom_httpx/test_http_handler.py | 0 .../custom_httpx/test_llm_http_handler.py | 0 .../test_databricks_chat_transformation.py | 0 .../test_databricks_common_utils.py | 0 ...seek_audio_transcription_transformation.py | 0 .../test_featherless_chat_transformation.py | 0 .../test_fireworks_ai_chat_transformation.py | 0 .../test_gemini_realtime_transformation.py | 0 .../test_hosted_vllm_chat_transformation.py | 0 .../test_llamafile_chat_transformation.py | 138 ++++++++--- .../test_lm_studio_chat_transformation.py | 8 +- .../test_meta_llama_chat_transformation.py | 0 .../test_mistral_chat_transformation.py | 0 .../chat/test_novita_chat_transformation.py | 0 .../chat/test_nscale_chat_transformation.py | 0 .../ollama/test_ollama_chat_transformation.py | 0 .../test_ollama_completion_transformation.py | 77 +++--- .../llms/ollama/test_ollama_model_info.py | 56 +++-- .../test_openai_responses_transformation.py | 0 .../openai/test_o_series_transformation.py | 12 +- .../llms/openai/test_openai_common_utils.py | 1 - .../test_openrouter_chat_transformation.py | 10 +- .../sagemaker/test_sagemaker_common_utils.py | 13 +- ...test_vertex_and_google_ai_studio_gemini.py | 0 ..._ai_multimodal_embedding_transformation.py | 0 .../llms/vertex_ai/test_http_status_201.py | 46 ++-- .../llms/vertex_ai/test_vertex.py | 28 +-- .../vertex_ai/test_vertex_ai_common_utils.py | 11 +- .../llms/vertex_ai/test_vertex_llm_base.py | 0 ...ai_partner_models_llama3_transformation.py | 0 tests/{litellm => test_litellm}/log.txt | 0 .../proxy/anthropic_endpoints/__init__.py | 0 .../anthropic_endpoints/test_endpoints.py | 49 ++-- .../proxy/auth/test_auth_checks.py | 0 .../proxy/auth/test_auth_exception_handler.py | 0 .../proxy/auth/test_handle_jwt.py | 0 .../proxy/auth/test_user_api_key_auth.py | 0 .../test_litellm/proxy/client/cli/__init__.py | 1 + .../proxy/client/cli/test_chat_commands.py | 4 +- .../client/cli/test_credentials_commands.py | 16 +- .../proxy/client/cli/test_global_options.py | 24 +- .../proxy/client/cli/test_keys_commands.py | 46 +++- .../proxy/client/cli/test_models_commands.py | 89 +++++-- .../proxy/client/cli/test_users_commands.py | 53 ++++- .../proxy/client/test_chat.py | 40 +++- .../proxy/client/test_client.py | 5 +- .../proxy/client/test_credentials.py | 59 ++++- .../proxy/client/test_http_client.py | 2 + .../proxy/client/test_http_commands.py | 3 +- .../proxy/client/test_keys.py | 52 ++++- .../proxy/client/test_model_groups.py | 23 +- .../proxy/client/test_models.py | 180 +++++++++++--- .../proxy/client/test_users.py | 21 +- .../common_utils/test_http_parsing_utils.py | 0 .../common_utils/test_reset_budget_job.py | 2 +- .../proxy/common_utils/test_timezone_utils.py | 0 .../test_base_update_queue.py | 0 .../test_daily_spend_update_queue.py | 0 .../test_pod_lock_manager.py | 0 .../test_spend_update_queue.py | 0 .../proxy/db/test_check_migration.py | 0 .../proxy/db/test_db_spend_update_writer.py | 0 .../proxy/db/test_exception_handler.py | 0 .../proxy/db/test_prisma_client.py | 0 .../mcp_server/test_tool_registry.py | 0 .../guardrails/test_guardrail_endpoints.py | 0 .../proxy/guardrails/test_init_guardrails.py | 0 .../health_endpoints/test_health_endpoints.py | 2 - .../hooks/test_parallel_request_limiter_v2.py | 0 .../hooks/test_proxy_track_cost_callback.py | 0 .../scim/test_scim_transformations.py | 0 .../test_common_daily_activity.py | 0 .../test_customer_endpoints.py | 0 .../test_internal_user_endpoints.py | 0 .../test_key_management_endpoints.py | 0 .../test_model_management_endpoints.py | 36 +-- .../test_tag_management_endpoints.py | 0 .../test_team_endpoints.py | 0 .../proxy/management_endpoints/test_ui_sso.py | 0 .../test_prometheus_auth_middleware.py | 0 .../test_files_endpoint.py | 0 .../test_llm_pass_through_endpoints.py | 0 .../test_pass_through_endpoints.py | 0 ...test_passthrough_endpoints_common_utils.py | 0 .../test_spend_management_endpoints.py | 0 .../test_spend_tracking_utils.py | 0 .../proxy/test_caching_routes.py | 0 .../proxy/test_common_request_processing.py | 39 ++-- .../test_configs/test_config_no_auth.yaml | 0 .../proxy/test_litellm_pre_call_utils.py | 20 +- .../proxy/test_proxy_cli.py | 54 ++--- .../proxy/test_proxy_server.py | 2 + .../proxy/test_route_llm_request.py | 32 +-- .../proxy/test_spend_log_cleanup.py | 35 ++- .../proxy/test_team_member_update.py | 9 +- .../test_litellm_proxy_types_utils.py | 0 .../test_proxy_setting_endpoints.py | 1 - tests/{litellm => test_litellm}/readme.md | 0 .../responses/test_responses_utils.py | 0 .../test_base_routing_strategy.py | 0 .../test_responses_api_deployment_check.py | 0 .../test_get_azure_ad_token_provider.py | 0 .../test_constants.py | 0 .../test_cost_calculator.py | 0 .../{litellm => test_litellm}/test_logging.py | 0 tests/{litellm => test_litellm}/test_main.py | 0 .../{litellm => test_litellm}/test_router.py | 5 +- tests/{litellm => test_litellm}/test_utils.py | 0 .../types/llms/test_types_llms_openai.py | 0 170 files changed, 1371 insertions(+), 793 deletions(-) rename tests/{litellm/proxy/anthropic_endpoints => }/__init__.py (100%) delete mode 100644 tests/litellm/integrations/test_athina.py delete mode 100644 tests/litellm/proxy/client/cli/__init__.py create mode 100644 tests/test_litellm/__init__.py rename tests/{litellm => test_litellm}/caching/test_in_memory_cache.py (100%) rename tests/{litellm => test_litellm}/caching/test_redis_cache.py (96%) rename tests/{litellm => test_litellm}/caching/test_redis_cluster_cache.py (100%) rename tests/{litellm => test_litellm}/caching/test_redis_semantic_cache.py (73%) rename tests/{litellm => test_litellm}/conftest.py (100%) rename tests/{litellm => test_litellm}/enterprise/enterprise_callbacks/send_emails/test_base_email.py (100%) rename tests/{litellm => test_litellm}/enterprise/enterprise_callbacks/send_emails/test_endpoints.py (100%) rename tests/{litellm => test_litellm}/enterprise/enterprise_callbacks/send_emails/test_resend_email.py (100%) rename tests/{litellm => test_litellm}/experimental_mcp_client/test_tools.py (100%) rename tests/{litellm => test_litellm}/integrations/SlackAlerting/test_slack_alerting.py (97%) rename tests/{litellm => test_litellm}/integrations/SlackAlerting/test_slack_alerting_utils.py (100%) rename tests/{litellm => test_litellm}/integrations/arize/test_arize_phoenix.py (99%) rename tests/{litellm => test_litellm}/integrations/arize/test_arize_utils.py (99%) rename tests/{litellm => test_litellm}/integrations/gcs_pubsub/test_pub_sub.py (100%) rename tests/{litellm => test_litellm}/integrations/test_agentops.py (70%) rename tests/{litellm => test_litellm}/integrations/test_anthropic_cache_control_hook.py (100%) create mode 100644 tests/test_litellm/integrations/test_athina.py rename tests/{litellm => test_litellm}/integrations/test_custom_prompt_management.py (100%) rename tests/{litellm => test_litellm}/integrations/test_deepeval.py (95%) rename tests/{litellm => test_litellm}/integrations/test_langfuse.py (100%) rename tests/{litellm => test_litellm}/integrations/test_opentelemetry.py (100%) rename tests/{litellm => test_litellm}/integrations/test_prometheus.py (100%) rename tests/{litellm => test_litellm}/integrations/test_prometheus_services.py (100%) create mode 100644 tests/test_litellm/litellm_core_utils/__init__.py rename tests/{litellm => test_litellm}/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py (100%) rename tests/{litellm => test_litellm}/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py (100%) rename tests/{litellm => test_litellm}/litellm_core_utils/messages_with_counts.py (99%) rename tests/{litellm => test_litellm}/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py (100%) rename tests/{litellm => test_litellm}/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py (97%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_core_helpers.py (100%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_dd_tracing.py (100%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_duration_parser.py (92%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_litellm_logging.py (100%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_realtime_streaming.py (100%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_safe_json_dumps.py (100%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_streaming_chunk_builder_utils.py (100%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_streaming_handler.py (97%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_token_counter.py (99%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_token_counter_tool.py (92%) rename tests/{litellm => test_litellm}/litellm_core_utils/test_token_counter_tool_data.py (86%) rename tests/{litellm => test_litellm}/llms/anthropic/chat/test_anthropic_chat_handler.py (100%) rename tests/{litellm => test_litellm}/llms/anthropic/chat/test_anthropic_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/azure/image_generation/test_azure_image_generation_init.py (100%) rename tests/{litellm => test_litellm}/llms/azure/test_azure_common_utils.py (100%) rename tests/{litellm => test_litellm}/llms/azure_ai/chat/test_azure_ai_transformation.py (68%) rename tests/{litellm => test_litellm}/llms/bedrock/chat/test_converse_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/bedrock/chat/test_invoke_handler.py (100%) rename tests/{litellm => test_litellm}/llms/bedrock/chat/test_mistral_config.py (85%) rename tests/{litellm => test_litellm}/llms/bedrock/image/test_amazon_stability3_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/bedrock/rerank/transformation.py (100%) rename tests/{litellm => test_litellm}/llms/bedrock/test_base_aws_llm.py (100%) rename tests/{litellm => test_litellm}/llms/bedrock/test_bedrock_common_utils.py (100%) rename tests/{litellm => test_litellm}/llms/chat/test_converse_handler.py (100%) rename tests/{litellm => test_litellm}/llms/cohere/chat/test_transformation.py (85%) rename tests/{litellm => test_litellm}/llms/custom_httpx/test_http_handler.py (100%) rename tests/{litellm => test_litellm}/llms/custom_httpx/test_llm_http_handler.py (100%) rename tests/{litellm => test_litellm}/llms/databricks/chat/test_databricks_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/databricks/test_databricks_common_utils.py (100%) rename tests/{litellm => test_litellm}/llms/deepgram/audio_transcription/test_deepseek_audio_transcription_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/featherless_ai/chat/test_featherless_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/gemini/realtime/test_gemini_realtime_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/llamafile/chat/test_llamafile_chat_transformation.py (52%) rename tests/{litellm => test_litellm}/llms/lm_studio/test_lm_studio_chat_transformation.py (88%) rename tests/{litellm => test_litellm}/llms/meta_llama/test_meta_llama_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/mistral/test_mistral_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/novita/chat/test_novita_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/nscale/chat/test_nscale_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/ollama/test_ollama_chat_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/ollama/test_ollama_completion_transformation.py (81%) rename tests/{litellm => test_litellm}/llms/ollama/test_ollama_model_info.py (71%) rename tests/{litellm => test_litellm}/llms/openai/responses/test_openai_responses_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/openai/test_o_series_transformation.py (86%) rename tests/{litellm => test_litellm}/llms/openai/test_openai_common_utils.py (99%) rename tests/{litellm => test_litellm}/llms/openrouter/chat/test_openrouter_chat_transformation.py (95%) rename tests/{litellm => test_litellm}/llms/sagemaker/test_sagemaker_common_utils.py (95%) rename tests/{litellm => test_litellm}/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py (100%) rename tests/{litellm => test_litellm}/llms/vertex_ai/multimodal_embeddings/test_vertex_ai_multimodal_embedding_transformation.py (100%) rename tests/{litellm => test_litellm}/llms/vertex_ai/test_http_status_201.py (90%) rename tests/{litellm => test_litellm}/llms/vertex_ai/test_vertex.py (99%) rename tests/{litellm => test_litellm}/llms/vertex_ai/test_vertex_ai_common_utils.py (99%) rename tests/{litellm => test_litellm}/llms/vertex_ai/test_vertex_llm_base.py (100%) rename tests/{litellm => test_litellm}/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py (100%) rename tests/{litellm => test_litellm}/log.txt (100%) create mode 100644 tests/test_litellm/proxy/anthropic_endpoints/__init__.py rename tests/{litellm => test_litellm}/proxy/anthropic_endpoints/test_endpoints.py (65%) rename tests/{litellm => test_litellm}/proxy/auth/test_auth_checks.py (100%) rename tests/{litellm => test_litellm}/proxy/auth/test_auth_exception_handler.py (100%) rename tests/{litellm => test_litellm}/proxy/auth/test_handle_jwt.py (100%) rename tests/{litellm => test_litellm}/proxy/auth/test_user_api_key_auth.py (100%) create mode 100644 tests/test_litellm/proxy/client/cli/__init__.py rename tests/{litellm => test_litellm}/proxy/client/cli/test_chat_commands.py (99%) rename tests/{litellm => test_litellm}/proxy/client/cli/test_credentials_commands.py (94%) rename tests/{litellm => test_litellm}/proxy/client/cli/test_global_options.py (74%) rename tests/{litellm => test_litellm}/proxy/client/cli/test_keys_commands.py (71%) rename tests/{litellm => test_litellm}/proxy/client/cli/test_models_commands.py (84%) rename tests/{litellm => test_litellm}/proxy/client/cli/test_users_commands.py (62%) rename tests/{litellm => test_litellm}/proxy/client/test_chat.py (81%) rename tests/{litellm => test_litellm}/proxy/client/test_client.py (97%) rename tests/{litellm => test_litellm}/proxy/client/test_credentials.py (78%) rename tests/{litellm => test_litellm}/proxy/client/test_http_client.py (99%) rename tests/{litellm => test_litellm}/proxy/client/test_http_commands.py (98%) rename tests/{litellm => test_litellm}/proxy/client/test_keys.py (88%) rename tests/{litellm => test_litellm}/proxy/client/test_model_groups.py (88%) rename tests/{litellm => test_litellm}/proxy/client/test_models.py (81%) rename tests/{litellm => test_litellm}/proxy/client/test_users.py (91%) rename tests/{litellm => test_litellm}/proxy/common_utils/test_http_parsing_utils.py (100%) rename tests/{litellm => test_litellm}/proxy/common_utils/test_reset_budget_job.py (99%) rename tests/{litellm => test_litellm}/proxy/common_utils/test_timezone_utils.py (100%) rename tests/{litellm => test_litellm}/proxy/db/db_transaction_queue/test_base_update_queue.py (100%) rename tests/{litellm => test_litellm}/proxy/db/db_transaction_queue/test_daily_spend_update_queue.py (100%) rename tests/{litellm => test_litellm}/proxy/db/db_transaction_queue/test_pod_lock_manager.py (100%) rename tests/{litellm => test_litellm}/proxy/db/db_transaction_queue/test_spend_update_queue.py (100%) rename tests/{litellm => test_litellm}/proxy/db/test_check_migration.py (100%) rename tests/{litellm => test_litellm}/proxy/db/test_db_spend_update_writer.py (100%) rename tests/{litellm => test_litellm}/proxy/db/test_exception_handler.py (100%) rename tests/{litellm => test_litellm}/proxy/db/test_prisma_client.py (100%) rename tests/{litellm => test_litellm}/proxy/experimental/mcp_server/test_tool_registry.py (100%) rename tests/{litellm => test_litellm}/proxy/guardrails/test_guardrail_endpoints.py (100%) rename tests/{litellm => test_litellm}/proxy/guardrails/test_init_guardrails.py (100%) rename tests/{litellm => test_litellm}/proxy/health_endpoints/test_health_endpoints.py (99%) rename tests/{litellm => test_litellm}/proxy/hooks/test_parallel_request_limiter_v2.py (100%) rename tests/{litellm => test_litellm}/proxy/hooks/test_proxy_track_cost_callback.py (100%) rename tests/{litellm => test_litellm}/proxy/management_endpoints/scim/test_scim_transformations.py (100%) rename tests/{litellm => test_litellm}/proxy/management_endpoints/test_common_daily_activity.py (100%) rename tests/{litellm => test_litellm}/proxy/management_endpoints/test_customer_endpoints.py (100%) rename tests/{litellm => test_litellm}/proxy/management_endpoints/test_internal_user_endpoints.py (100%) rename tests/{litellm => test_litellm}/proxy/management_endpoints/test_key_management_endpoints.py (100%) rename tests/{litellm => test_litellm}/proxy/management_endpoints/test_model_management_endpoints.py (96%) rename tests/{litellm => test_litellm}/proxy/management_endpoints/test_tag_management_endpoints.py (100%) rename tests/{litellm => test_litellm}/proxy/management_endpoints/test_team_endpoints.py (100%) rename tests/{litellm => test_litellm}/proxy/management_endpoints/test_ui_sso.py (100%) rename tests/{litellm => test_litellm}/proxy/middleware/test_prometheus_auth_middleware.py (100%) rename tests/{litellm => test_litellm}/proxy/openai_files_endpoint/test_files_endpoint.py (100%) rename tests/{litellm => test_litellm}/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py (100%) rename tests/{litellm => test_litellm}/proxy/pass_through_endpoints/test_pass_through_endpoints.py (100%) rename tests/{litellm => test_litellm}/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py (100%) rename tests/{litellm => test_litellm}/proxy/spend_tracking/test_spend_management_endpoints.py (100%) rename tests/{litellm => test_litellm}/proxy/spend_tracking/test_spend_tracking_utils.py (100%) rename tests/{litellm => test_litellm}/proxy/test_caching_routes.py (100%) rename tests/{litellm => test_litellm}/proxy/test_common_request_processing.py (93%) rename tests/{litellm => test_litellm}/proxy/test_configs/test_config_no_auth.yaml (100%) rename tests/{litellm => test_litellm}/proxy/test_litellm_pre_call_utils.py (95%) rename tests/{litellm => test_litellm}/proxy/test_proxy_cli.py (91%) rename tests/{litellm => test_litellm}/proxy/test_proxy_server.py (99%) rename tests/{litellm => test_litellm}/proxy/test_route_llm_request.py (87%) rename tests/{litellm => test_litellm}/proxy/test_spend_log_cleanup.py (85%) rename tests/{litellm => test_litellm}/proxy/test_team_member_update.py (92%) rename tests/{litellm => test_litellm}/proxy/types_utils/test_litellm_proxy_types_utils.py (100%) rename tests/{litellm => test_litellm}/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py (99%) rename tests/{litellm => test_litellm}/readme.md (100%) rename tests/{litellm => test_litellm}/responses/test_responses_utils.py (100%) rename tests/{litellm => test_litellm}/router_strategy/test_base_routing_strategy.py (100%) rename tests/{litellm => test_litellm}/router_utils/pre_call_checks/test_responses_api_deployment_check.py (100%) rename tests/{litellm => test_litellm}/secret_managers/test_get_azure_ad_token_provider.py (100%) rename tests/{litellm => test_litellm}/test_constants.py (100%) rename tests/{litellm => test_litellm}/test_cost_calculator.py (100%) rename tests/{litellm => test_litellm}/test_logging.py (100%) rename tests/{litellm => test_litellm}/test_main.py (100%) rename tests/{litellm => test_litellm}/test_router.py (99%) rename tests/{litellm => test_litellm}/test_utils.py (100%) rename tests/{litellm => test_litellm}/types/llms/test_types_llms_openai.py (100%) diff --git a/.github/workflows/test-litellm.yml b/.github/workflows/test-litellm.yml index a2b9e6c7c34..db2ebb3cbae 100644 --- a/.github/workflows/test-litellm.yml +++ b/.github/workflows/test-litellm.yml @@ -28,6 +28,7 @@ jobs: - name: Install dependencies run: | poetry install --with dev,proxy-dev --extras proxy + poetry run pip install "pytest-retry==1.6.3" poetry run pip install pytest-xdist - name: Setup litellm-enterprise as local package run: | @@ -36,4 +37,4 @@ jobs: cd .. - name: Run tests run: | - poetry run pytest tests/litellm -x -vv -n 4 \ No newline at end of file + poetry run pytest tests/test_litellm -x -vv -n 4 \ No newline at end of file diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 9143727163c..dd98498e3be 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -24,7 +24,7 @@ repos: rev: 7.0.0 # The version of flake8 to use hooks: - id: flake8 - exclude: ^litellm/tests/|^litellm/proxy/tests/|^litellm/tests/litellm/|^tests/litellm/ + exclude: ^litellm/tests/|^litellm/proxy/tests/|^litellm/tests/test_litellm/|^tests/test_litellm/ additional_dependencies: [flake8-print] files: (litellm/|litellm_proxy_extras/|enterprise/).*\.py - repo: https://github.com/python-poetry/poetry diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index c736b461b18..0b679619747 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -4392,7 +4392,7 @@ "supports_function_calling": true, "supports_response_schema": true, "supports_tool_choice": true, - "deprecation_date": "2025-1-6" + "deprecation_date": "2025-01-06" }, "groq/llama3-groq-8b-8192-tool-use-preview": { "max_tokens": 8192, @@ -4405,7 +4405,7 @@ "supports_function_calling": true, "supports_response_schema": true, "supports_tool_choice": true, - "deprecation_date": "2025-1-6" + "deprecation_date": "2025-01-06" }, "groq/qwen-qwq-32b": { "max_tokens": 128000, diff --git a/tests/litellm/proxy/anthropic_endpoints/__init__.py b/tests/__init__.py similarity index 100% rename from tests/litellm/proxy/anthropic_endpoints/__init__.py rename to tests/__init__.py diff --git a/tests/litellm/integrations/test_athina.py b/tests/litellm/integrations/test_athina.py deleted file mode 100644 index fd660a036ed..00000000000 --- a/tests/litellm/integrations/test_athina.py +++ /dev/null @@ -1,207 +0,0 @@ -import unittest -from unittest.mock import patch, MagicMock, ANY -import json -import datetime -import sys -import os - -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system-path - -from litellm.integrations.athina import AthinaLogger - -class TestAthinaLogger(unittest.TestCase): - - def setUp(self): - # Set up environment variables for testing - self.env_patcher = patch.dict('os.environ', { - 'ATHINA_API_KEY': 'test-api-key', - 'ATHINA_BASE_URL': 'https://test.athina.ai' - }) - self.env_patcher.start() - self.logger = AthinaLogger() - - # Setup common test variables - self.start_time = datetime.datetime(2023, 1, 1, 12, 0, 0) - self.end_time = datetime.datetime(2023, 1, 1, 12, 0, 1) - self.print_verbose = MagicMock() - - def tearDown(self): - self.env_patcher.stop() - - def test_init(self): - """Test the initialization of AthinaLogger""" - self.assertEqual(self.logger.athina_api_key, 'test-api-key') - self.assertEqual(self.logger.athina_logging_url, 'https://test.athina.ai/api/v1/log/inference') - self.assertEqual(self.logger.headers, { - 'athina-api-key': 'test-api-key', - 'Content-Type': 'application/json' - }) - - @patch('litellm.module_level_client.post') - def test_log_event_success(self, mock_post): - """Test successful logging of an event""" - # Setup mock response - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.text = "Success" - mock_post.return_value = mock_response - - # Create test data - kwargs = { - 'model': 'gpt-4', - 'messages': [{'role': 'user', 'content': 'Hello'}], - 'stream': False, - 'litellm_params': { - 'metadata': { - 'environment': 'test-environment', - 'prompt_slug': 'test-prompt', - 'customer_id': 'test-customer', - 'customer_user_id': 'test-user', - 'session_id': 'test-session', - 'external_reference_id': 'test-ext-ref', - 'context': 'test-context', - 'expected_response': 'test-expected', - 'user_query': 'test-query', - 'tags': ['test-tag'], - 'user_feedback': 'test-feedback', - 'model_options': {'test-opt': 'test-val'}, - 'custom_attributes': {'test-attr': 'test-val'} - } - } - } - - response_obj = MagicMock() - response_obj.model_dump.return_value = { - 'id': 'resp-123', - 'choices': [{'message': {'content': 'Hi there'}}], - 'usage': { - 'prompt_tokens': 10, - 'completion_tokens': 5, - 'total_tokens': 15 - } - } - - # Call the method - self.logger.log_event(kwargs, response_obj, self.start_time, self.end_time, self.print_verbose) - - # Verify the results - mock_post.assert_called_once() - call_args = mock_post.call_args - self.assertEqual(call_args[0][0], 'https://test.athina.ai/api/v1/log/inference') - self.assertEqual(call_args[1]['headers'], self.logger.headers) - - # Parse and verify the sent data - sent_data = json.loads(call_args[1]['data']) - self.assertEqual(sent_data['language_model_id'], 'gpt-4') - self.assertEqual(sent_data['prompt'], kwargs['messages']) - self.assertEqual(sent_data['prompt_tokens'], 10) - self.assertEqual(sent_data['completion_tokens'], 5) - self.assertEqual(sent_data['total_tokens'], 15) - self.assertEqual(sent_data['response_time'], 1000) # 1 second = 1000ms - self.assertEqual(sent_data['customer_id'], 'test-customer') - self.assertEqual(sent_data['session_id'], 'test-session') - self.assertEqual(sent_data['environment'], 'test-environment') - self.assertEqual(sent_data['prompt_slug'], 'test-prompt') - self.assertEqual(sent_data['external_reference_id'], 'test-ext-ref') - self.assertEqual(sent_data['context'], 'test-context') - self.assertEqual(sent_data['expected_response'], 'test-expected') - self.assertEqual(sent_data['user_query'], 'test-query') - self.assertEqual(sent_data['tags'], ['test-tag']) - self.assertEqual(sent_data['user_feedback'], 'test-feedback') - self.assertEqual(sent_data['model_options'], {'test-opt': 'test-val'}) - self.assertEqual(sent_data['custom_attributes'], {'test-attr': 'test-val'}) - # Verify the print_verbose was called - self.print_verbose.assert_called_once_with("Athina Logger Succeeded - Success") - - @patch('litellm.module_level_client.post') - def test_log_event_error_response(self, mock_post): - """Test handling of error response from the API""" - # Setup mock error response - mock_response = MagicMock() - mock_response.status_code = 400 - mock_response.text = "Bad Request" - mock_post.return_value = mock_response - - # Create test data - kwargs = { - 'model': 'gpt-4', - 'messages': [{'role': 'user', 'content': 'Hello'}], - 'stream': False - } - - response_obj = MagicMock() - response_obj.model_dump.return_value = { - 'id': 'resp-123', - 'choices': [{'message': {'content': 'Hi there'}}], - 'usage': { - 'prompt_tokens': 10, - 'completion_tokens': 5, - 'total_tokens': 15 - } - } - - # Call the method - self.logger.log_event(kwargs, response_obj, self.start_time, self.end_time, self.print_verbose) - - # Verify print_verbose was called with error message - self.print_verbose.assert_called_once_with("Athina Logger Error - Bad Request, 400") - - @patch('litellm.module_level_client.post') - def test_log_event_exception(self, mock_post): - """Test handling of exceptions during logging""" - # Setup mock to raise exception - mock_post.side_effect = Exception("Test exception") - - # Create test data - kwargs = { - 'model': 'gpt-4', - 'messages': [{'role': 'user', 'content': 'Hello'}], - 'stream': False - } - - response_obj = MagicMock() - response_obj.model_dump.return_value = {} - - # Call the method - self.logger.log_event(kwargs, response_obj, self.start_time, self.end_time, self.print_verbose) - - # Verify print_verbose was called with exception info - self.print_verbose.assert_called_once() - self.assertIn("Athina Logger Error - Test exception", self.print_verbose.call_args[0][0]) - - @patch('litellm.module_level_client.post') - def test_log_event_with_tools(self, mock_post): - """Test logging with tools/functions data""" - # Setup mock response - mock_response = MagicMock() - mock_response.status_code = 200 - mock_post.return_value = mock_response - - # Create test data with tools - kwargs = { - 'model': 'gpt-4', - 'messages': [{'role': 'user', 'content': "What's the weather?"}], - 'stream': False, - 'optional_params': { - 'tools': [{'type': 'function', 'function': {'name': 'get_weather'}}] - } - } - - response_obj = MagicMock() - response_obj.model_dump.return_value = { - 'id': 'resp-123', - 'usage': {'prompt_tokens': 10, 'completion_tokens': 5, 'total_tokens': 15} - } - - # Call the method - self.logger.log_event(kwargs, response_obj, self.start_time, self.end_time, self.print_verbose) - - # Verify the results - sent_data = json.loads(mock_post.call_args[1]['data']) - self.assertEqual(sent_data['tools'], [{'type': 'function', 'function': {'name': 'get_weather'}}]) - - -if __name__ == '__main__': - unittest.main() \ No newline at end of file diff --git a/tests/litellm/proxy/client/cli/__init__.py b/tests/litellm/proxy/client/cli/__init__.py deleted file mode 100644 index 352d47cd995..00000000000 --- a/tests/litellm/proxy/client/cli/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Tests for the LiteLLM Proxy Client CLI package.""" \ No newline at end of file diff --git a/tests/test_litellm/__init__.py b/tests/test_litellm/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/litellm/caching/test_in_memory_cache.py b/tests/test_litellm/caching/test_in_memory_cache.py similarity index 100% rename from tests/litellm/caching/test_in_memory_cache.py rename to tests/test_litellm/caching/test_in_memory_cache.py diff --git a/tests/litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py similarity index 96% rename from tests/litellm/caching/test_redis_cache.py rename to tests/test_litellm/caching/test_redis_cache.py index 87502427067..447b0a3bbd3 100644 --- a/tests/litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -16,7 +16,7 @@ from litellm.caching.redis_cache import RedisCache @pytest.fixture def redis_no_ping(): """Patch RedisCache initialization to prevent async ping tasks from being created""" - with patch('asyncio.get_running_loop') as mock_get_loop: + with patch("asyncio.get_running_loop") as mock_get_loop: # Either raise an exception or return a mock that will handle the task creation mock_get_loop.side_effect = RuntimeError("No running event loop") yield @@ -64,32 +64,32 @@ async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping) async def test_redis_cache_async_batch_get_cache(monkeypatch, redis_no_ping): monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache() - + # Create an AsyncMock for the Redis client mock_redis_instance = AsyncMock() - + # Make sure the mock can be used as an async context manager mock_redis_instance.__aenter__.return_value = mock_redis_instance mock_redis_instance.__aexit__.return_value = None - + # Setup the return value for mget mock_redis_instance.mget.return_value = [ b'{"key1": "value1"}', None, - b'{"key3": "value3"}' + b'{"key3": "value3"}', ] - + test_keys = ["key1", "key2", "key3"] - + with patch.object( redis_cache, "init_async_client", return_value=mock_redis_instance ): # Call async_batch_get_cache result = await redis_cache.async_batch_get_cache(key_list=test_keys) - + # Verify mget was called with the correct keys mock_redis_instance.mget.assert_called_once() - + # Check that results were properly decoded assert result["key1"] == {"key1": "value1"} assert result["key2"] is None diff --git a/tests/litellm/caching/test_redis_cluster_cache.py b/tests/test_litellm/caching/test_redis_cluster_cache.py similarity index 100% rename from tests/litellm/caching/test_redis_cluster_cache.py rename to tests/test_litellm/caching/test_redis_cluster_cache.py diff --git a/tests/litellm/caching/test_redis_semantic_cache.py b/tests/test_litellm/caching/test_redis_semantic_cache.py similarity index 73% rename from tests/litellm/caching/test_redis_semantic_cache.py rename to tests/test_litellm/caching/test_redis_semantic_cache.py index 142f7990c42..f9946e266fe 100644 --- a/tests/litellm/caching/test_redis_semantic_cache.py +++ b/tests/test_litellm/caching/test_redis_semantic_cache.py @@ -1,6 +1,6 @@ import os import sys -from unittest.mock import MagicMock, patch, AsyncMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -13,27 +13,30 @@ sys.path.insert( def test_redis_semantic_cache_initialization(monkeypatch): # Mock the redisvl import semantic_cache_mock = MagicMock() - with patch.dict("sys.modules", { - "redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock), - "redisvl.utils.vectorize": MagicMock(CustomTextVectorizer=MagicMock()) - }): + with patch.dict( + "sys.modules", + { + "redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock), + "redisvl.utils.vectorize": MagicMock(CustomTextVectorizer=MagicMock()), + }, + ): from litellm.caching.redis_semantic_cache import RedisSemanticCache - + # Set environment variables monkeypatch.setenv("REDIS_HOST", "localhost") monkeypatch.setenv("REDIS_PORT", "6379") monkeypatch.setenv("REDIS_PASSWORD", "test_password") - + # Initialize the cache with a similarity threshold redis_semantic_cache = RedisSemanticCache(similarity_threshold=0.8) - + # Verify the semantic cache was initialized with correct parameters assert redis_semantic_cache.similarity_threshold == 0.8 - + # Use pytest.approx for floating point comparison to handle precision issues assert redis_semantic_cache.distance_threshold == pytest.approx(0.2, abs=1e-10) assert redis_semantic_cache.embedding_model == "text-embedding-ada-002" - + # Test initialization with missing similarity_threshold with pytest.raises(ValueError, match="similarity_threshold must be provided"): RedisSemanticCache() @@ -43,42 +46,48 @@ def test_redis_semantic_cache_get_cache(monkeypatch): # Mock the redisvl import and embedding function semantic_cache_mock = MagicMock() custom_vectorizer_mock = MagicMock() - - with patch.dict("sys.modules", { - "redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock), - "redisvl.utils.vectorize": MagicMock(CustomTextVectorizer=custom_vectorizer_mock) - }): + + with patch.dict( + "sys.modules", + { + "redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock), + "redisvl.utils.vectorize": MagicMock( + CustomTextVectorizer=custom_vectorizer_mock + ), + }, + ): from litellm.caching.redis_semantic_cache import RedisSemanticCache - + # Set environment variables monkeypatch.setenv("REDIS_HOST", "localhost") monkeypatch.setenv("REDIS_PORT", "6379") monkeypatch.setenv("REDIS_PASSWORD", "test_password") - + # Initialize cache redis_semantic_cache = RedisSemanticCache(similarity_threshold=0.8) - + # Mock the llmcache.check method to return a result mock_result = [ { "prompt": "What is the capital of France?", "response": '{"content": "Paris is the capital of France."}', - "vector_distance": 0.1 # Distance of 0.1 means similarity of 0.9 + "vector_distance": 0.1, # Distance of 0.1 means similarity of 0.9 } ] redis_semantic_cache.llmcache.check = MagicMock(return_value=mock_result) - + # Mock the embedding function - with patch("litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2, 0.3]}]}): + with patch( + "litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2, 0.3]}]} + ): # Test get_cache with a message result = redis_semantic_cache.get_cache( - key="test_key", - messages=[{"content": "What is the capital of France?"}] + key="test_key", messages=[{"content": "What is the capital of France?"}] ) - + # Verify result is properly parsed assert result == {"content": "Paris is the capital of France."} - + # Verify llmcache.check was called redis_semantic_cache.llmcache.check.assert_called_once() @@ -88,43 +97,50 @@ async def test_redis_semantic_cache_async_get_cache(monkeypatch): # Mock the redisvl import semantic_cache_mock = MagicMock() custom_vectorizer_mock = MagicMock() - - with patch.dict("sys.modules", { - "redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock), - "redisvl.utils.vectorize": MagicMock(CustomTextVectorizer=custom_vectorizer_mock) - }): + + with patch.dict( + "sys.modules", + { + "redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock), + "redisvl.utils.vectorize": MagicMock( + CustomTextVectorizer=custom_vectorizer_mock + ), + }, + ): from litellm.caching.redis_semantic_cache import RedisSemanticCache - + # Set environment variables monkeypatch.setenv("REDIS_HOST", "localhost") monkeypatch.setenv("REDIS_PORT", "6379") monkeypatch.setenv("REDIS_PASSWORD", "test_password") - + # Initialize cache redis_semantic_cache = RedisSemanticCache(similarity_threshold=0.8) - + # Mock the async methods mock_result = [ { "prompt": "What is the capital of France?", "response": '{"content": "Paris is the capital of France."}', - "vector_distance": 0.1 # Distance of 0.1 means similarity of 0.9 + "vector_distance": 0.1, # Distance of 0.1 means similarity of 0.9 } ] - + redis_semantic_cache.llmcache.acheck = AsyncMock(return_value=mock_result) - redis_semantic_cache._get_async_embedding = AsyncMock(return_value=[0.1, 0.2, 0.3]) - + redis_semantic_cache._get_async_embedding = AsyncMock( + return_value=[0.1, 0.2, 0.3] + ) + # Test async_get_cache with a message result = await redis_semantic_cache.async_get_cache( key="test_key", messages=[{"content": "What is the capital of France?"}], - metadata={} + metadata={}, ) - + # Verify result is properly parsed assert result == {"content": "Paris is the capital of France."} - + # Verify methods were called redis_semantic_cache._get_async_embedding.assert_called_once() - redis_semantic_cache.llmcache.acheck.assert_called_once() \ No newline at end of file + redis_semantic_cache.llmcache.acheck.assert_called_once() diff --git a/tests/litellm/conftest.py b/tests/test_litellm/conftest.py similarity index 100% rename from tests/litellm/conftest.py rename to tests/test_litellm/conftest.py diff --git a/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py similarity index 100% rename from tests/litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py rename to tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py diff --git a/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py similarity index 100% rename from tests/litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py rename to tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py diff --git a/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py similarity index 100% rename from tests/litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py rename to tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py diff --git a/tests/litellm/experimental_mcp_client/test_tools.py b/tests/test_litellm/experimental_mcp_client/test_tools.py similarity index 100% rename from tests/litellm/experimental_mcp_client/test_tools.py rename to tests/test_litellm/experimental_mcp_client/test_tools.py diff --git a/tests/litellm/integrations/SlackAlerting/test_slack_alerting.py b/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py similarity index 97% rename from tests/litellm/integrations/SlackAlerting/test_slack_alerting.py rename to tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py index b9d0328c5f1..d389be79618 100644 --- a/tests/litellm/integrations/SlackAlerting/test_slack_alerting.py +++ b/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting.py @@ -162,13 +162,13 @@ class TestSlackAlerting(unittest.TestCase): self.assertEqual(event, "soft_budget_crossed") self.assertTrue("Total Soft Budget" in event_message) - # Calling update_values with alerting args should try to start the periodic task + # Calling update_values with alerting args should try to start the periodic task @patch("asyncio.create_task") def test_update_values_starts_periodic_task(self, mock_create_task): # Make it do nothing (or return a dummy future) mock_create_task.return_value = AsyncMock() # prevents awaiting errors - assert(self.slack_alerting.periodic_started == False) + assert self.slack_alerting.periodic_started == False self.slack_alerting.update_values(alerting_args={"slack_alerting": "True"}) - assert(self.slack_alerting.periodic_started == True) + assert self.slack_alerting.periodic_started == True diff --git a/tests/litellm/integrations/SlackAlerting/test_slack_alerting_utils.py b/tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_utils.py similarity index 100% rename from tests/litellm/integrations/SlackAlerting/test_slack_alerting_utils.py rename to tests/test_litellm/integrations/SlackAlerting/test_slack_alerting_utils.py diff --git a/tests/litellm/integrations/arize/test_arize_phoenix.py b/tests/test_litellm/integrations/arize/test_arize_phoenix.py similarity index 99% rename from tests/litellm/integrations/arize/test_arize_phoenix.py rename to tests/test_litellm/integrations/arize/test_arize_phoenix.py index 4f2fb3b9d53..fd81d9d9d7e 100644 --- a/tests/litellm/integrations/arize/test_arize_phoenix.py +++ b/tests/test_litellm/integrations/arize/test_arize_phoenix.py @@ -5,7 +5,6 @@ from litellm.integrations.arize.arize_phoenix import ArizePhoenixLogger class TestArizePhoenixConfig(unittest.TestCase): - @patch.dict( "os.environ", { diff --git a/tests/litellm/integrations/arize/test_arize_utils.py b/tests/test_litellm/integrations/arize/test_arize_utils.py similarity index 99% rename from tests/litellm/integrations/arize/test_arize_utils.py rename to tests/test_litellm/integrations/arize/test_arize_utils.py index bea42faaa8c..4286398aca0 100644 --- a/tests/litellm/integrations/arize/test_arize_utils.py +++ b/tests/test_litellm/integrations/arize/test_arize_utils.py @@ -7,15 +7,17 @@ from typing import Optional sys.path.insert(0, os.path.abspath("../..")) import asyncio -import litellm + import pytest -from litellm.integrations.arize.arize import ArizeLogger -from litellm.integrations.custom_logger import CustomLogger + +import litellm from litellm.integrations._types.open_inference import ( - SpanAttributes, MessageAttributes, + SpanAttributes, ToolCallAttributes, ) +from litellm.integrations.arize.arize import ArizeLogger +from litellm.integrations.custom_logger import CustomLogger from litellm.types.utils import Choices, StandardCallbackDynamicParams @@ -25,6 +27,7 @@ def test_arize_set_attributes(): Ensures that the correct span attributes are being added during a request. """ from unittest.mock import MagicMock + from litellm.types.utils import ModelResponse span = MagicMock() # Mocked tracing span to test attribute setting diff --git a/tests/litellm/integrations/gcs_pubsub/test_pub_sub.py b/tests/test_litellm/integrations/gcs_pubsub/test_pub_sub.py similarity index 100% rename from tests/litellm/integrations/gcs_pubsub/test_pub_sub.py rename to tests/test_litellm/integrations/gcs_pubsub/test_pub_sub.py diff --git a/tests/litellm/integrations/test_agentops.py b/tests/test_litellm/integrations/test_agentops.py similarity index 70% rename from tests/litellm/integrations/test_agentops.py rename to tests/test_litellm/integrations/test_agentops.py index 33027361773..85ee34a0d8c 100644 --- a/tests/litellm/integrations/test_agentops.py +++ b/tests/test_litellm/integrations/test_agentops.py @@ -1,18 +1,20 @@ import os import sys -import pytest -from unittest.mock import patch, MagicMock +from unittest.mock import MagicMock, patch -sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system-path +import pytest + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system-path from litellm.integrations.agentops.agentops import AgentOps, AgentOpsConfig + @pytest.fixture def mock_auth_response(): - return { - "token": "test_jwt_token", - "project_id": "test_project_id" - } + return {"token": "test_jwt_token", "project_id": "test_project_id"} + @pytest.fixture def agentops_config(): @@ -21,16 +23,20 @@ def agentops_config(): api_key="test_api_key", service_name="test_service", deployment_environment="test_env", - auth_endpoint="https://api.agentops.ai/v3/auth/token" + auth_endpoint="https://api.agentops.ai/v3/auth/token", ) + def test_agentops_config_from_env(): """Test that AgentOpsConfig correctly reads from environment variables""" - with patch.dict(os.environ, { - "AGENTOPS_API_KEY": "test_key", - "AGENTOPS_SERVICE_NAME": "test_service", - "AGENTOPS_ENVIRONMENT": "test_env" - }): + with patch.dict( + os.environ, + { + "AGENTOPS_API_KEY": "test_key", + "AGENTOPS_SERVICE_NAME": "test_service", + "AGENTOPS_ENVIRONMENT": "test_env", + }, + ): config = AgentOpsConfig.from_env() assert config.api_key == "test_key" assert config.service_name == "test_service" @@ -38,6 +44,7 @@ def test_agentops_config_from_env(): assert config.endpoint == "https://otlp.agentops.cloud/v1/traces" assert config.auth_endpoint == "https://api.agentops.ai/v3/auth/token" + def test_agentops_config_defaults(): """Test that AgentOpsConfig uses correct default values""" config = AgentOpsConfig() @@ -47,52 +54,64 @@ def test_agentops_config_defaults(): assert config.endpoint == "https://otlp.agentops.cloud/v1/traces" assert config.auth_endpoint == "https://api.agentops.ai/v3/auth/token" -@patch('litellm.integrations.agentops.agentops.AgentOps._fetch_auth_token') + +@patch("litellm.integrations.agentops.agentops.AgentOps._fetch_auth_token") def test_fetch_auth_token_success(mock_fetch_auth_token, mock_auth_response): """Test successful JWT token fetch""" mock_fetch_auth_token.return_value = mock_auth_response - + config = AgentOpsConfig(api_key="test_key") agentops = AgentOps(config=config) - - mock_fetch_auth_token.assert_called_once_with("test_key", "https://api.agentops.ai/v3/auth/token") - assert agentops.resource_attributes.get("project.id") == mock_auth_response.get("project_id") -@patch('litellm.integrations.agentops.agentops.AgentOps._fetch_auth_token') + mock_fetch_auth_token.assert_called_once_with( + "test_key", "https://api.agentops.ai/v3/auth/token" + ) + assert agentops.resource_attributes.get("project.id") == mock_auth_response.get( + "project_id" + ) + + +@patch("litellm.integrations.agentops.agentops.AgentOps._fetch_auth_token") def test_fetch_auth_token_failure(mock_fetch_auth_token): """Test failed JWT token fetch""" - mock_fetch_auth_token.side_effect = Exception("Failed to fetch auth token: Unauthorized") - + mock_fetch_auth_token.side_effect = Exception( + "Failed to fetch auth token: Unauthorized" + ) + config = AgentOpsConfig(api_key="test_key") agentops = AgentOps(config=config) - + mock_fetch_auth_token.assert_called_once() assert "project.id" not in agentops.resource_attributes -@patch('litellm.integrations.agentops.agentops.AgentOps._fetch_auth_token') -def test_agentops_initialization(mock_fetch_auth_token, agentops_config, mock_auth_response): + +@patch("litellm.integrations.agentops.agentops.AgentOps._fetch_auth_token") +def test_agentops_initialization( + mock_fetch_auth_token, agentops_config, mock_auth_response +): """Test AgentOps initialization with config""" mock_fetch_auth_token.return_value = mock_auth_response - + agentops = AgentOps(config=agentops_config) - + assert agentops.resource_attributes["service.name"] == "test_service" assert agentops.resource_attributes["deployment.environment"] == "test_env" assert agentops.resource_attributes["telemetry.sdk.name"] == "agentops" assert agentops.resource_attributes["project.id"] == "test_project_id" + def test_agentops_initialization_no_auth(): """Test AgentOps initialization without authentication""" test_config = AgentOpsConfig( endpoint="https://otlp.agentops.cloud/v1/traces", api_key=None, # No API key service_name="test_service", - deployment_environment="test_env" + deployment_environment="test_env", ) - + agentops = AgentOps(config=test_config) - + assert agentops.resource_attributes["service.name"] == "test_service" assert agentops.resource_attributes["deployment.environment"] == "test_env" assert agentops.resource_attributes["telemetry.sdk.name"] == "agentops" - assert "project.id" not in agentops.resource_attributes \ No newline at end of file + assert "project.id" not in agentops.resource_attributes diff --git a/tests/litellm/integrations/test_anthropic_cache_control_hook.py b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py similarity index 100% rename from tests/litellm/integrations/test_anthropic_cache_control_hook.py rename to tests/test_litellm/integrations/test_anthropic_cache_control_hook.py diff --git a/tests/test_litellm/integrations/test_athina.py b/tests/test_litellm/integrations/test_athina.py new file mode 100644 index 00000000000..49d8fc693e7 --- /dev/null +++ b/tests/test_litellm/integrations/test_athina.py @@ -0,0 +1,220 @@ +import datetime +import json +import os +import sys +import unittest +from unittest.mock import ANY, MagicMock, patch + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system-path + +from litellm.integrations.athina import AthinaLogger + + +class TestAthinaLogger(unittest.TestCase): + def setUp(self): + # Set up environment variables for testing + self.env_patcher = patch.dict( + "os.environ", + { + "ATHINA_API_KEY": "test-api-key", + "ATHINA_BASE_URL": "https://test.athina.ai", + }, + ) + self.env_patcher.start() + self.logger = AthinaLogger() + + # Setup common test variables + self.start_time = datetime.datetime(2023, 1, 1, 12, 0, 0) + self.end_time = datetime.datetime(2023, 1, 1, 12, 0, 1) + self.print_verbose = MagicMock() + + def tearDown(self): + self.env_patcher.stop() + + def test_init(self): + """Test the initialization of AthinaLogger""" + self.assertEqual(self.logger.athina_api_key, "test-api-key") + self.assertEqual( + self.logger.athina_logging_url, + "https://test.athina.ai/api/v1/log/inference", + ) + self.assertEqual( + self.logger.headers, + {"athina-api-key": "test-api-key", "Content-Type": "application/json"}, + ) + + @patch("litellm.module_level_client.post") + def test_log_event_success(self, mock_post): + """Test successful logging of an event""" + # Setup mock response + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.text = "Success" + mock_post.return_value = mock_response + + # Create test data + kwargs = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "stream": False, + "litellm_params": { + "metadata": { + "environment": "test-environment", + "prompt_slug": "test-prompt", + "customer_id": "test-customer", + "customer_user_id": "test-user", + "session_id": "test-session", + "external_reference_id": "test-ext-ref", + "context": "test-context", + "expected_response": "test-expected", + "user_query": "test-query", + "tags": ["test-tag"], + "user_feedback": "test-feedback", + "model_options": {"test-opt": "test-val"}, + "custom_attributes": {"test-attr": "test-val"}, + } + }, + } + + response_obj = MagicMock() + response_obj.model_dump.return_value = { + "id": "resp-123", + "choices": [{"message": {"content": "Hi there"}}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + + # Call the method + self.logger.log_event( + kwargs, response_obj, self.start_time, self.end_time, self.print_verbose + ) + + # Verify the results + mock_post.assert_called_once() + call_args = mock_post.call_args + self.assertEqual(call_args[0][0], "https://test.athina.ai/api/v1/log/inference") + self.assertEqual(call_args[1]["headers"], self.logger.headers) + + # Parse and verify the sent data + sent_data = json.loads(call_args[1]["data"]) + self.assertEqual(sent_data["language_model_id"], "gpt-4") + self.assertEqual(sent_data["prompt"], kwargs["messages"]) + self.assertEqual(sent_data["prompt_tokens"], 10) + self.assertEqual(sent_data["completion_tokens"], 5) + self.assertEqual(sent_data["total_tokens"], 15) + self.assertEqual(sent_data["response_time"], 1000) # 1 second = 1000ms + self.assertEqual(sent_data["customer_id"], "test-customer") + self.assertEqual(sent_data["session_id"], "test-session") + self.assertEqual(sent_data["environment"], "test-environment") + self.assertEqual(sent_data["prompt_slug"], "test-prompt") + self.assertEqual(sent_data["external_reference_id"], "test-ext-ref") + self.assertEqual(sent_data["context"], "test-context") + self.assertEqual(sent_data["expected_response"], "test-expected") + self.assertEqual(sent_data["user_query"], "test-query") + self.assertEqual(sent_data["tags"], ["test-tag"]) + self.assertEqual(sent_data["user_feedback"], "test-feedback") + self.assertEqual(sent_data["model_options"], {"test-opt": "test-val"}) + self.assertEqual(sent_data["custom_attributes"], {"test-attr": "test-val"}) + # Verify the print_verbose was called + self.print_verbose.assert_called_once_with("Athina Logger Succeeded - Success") + + @patch("litellm.module_level_client.post") + def test_log_event_error_response(self, mock_post): + """Test handling of error response from the API""" + # Setup mock error response + mock_response = MagicMock() + mock_response.status_code = 400 + mock_response.text = "Bad Request" + mock_post.return_value = mock_response + + # Create test data + kwargs = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "stream": False, + } + + response_obj = MagicMock() + response_obj.model_dump.return_value = { + "id": "resp-123", + "choices": [{"message": {"content": "Hi there"}}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + + # Call the method + self.logger.log_event( + kwargs, response_obj, self.start_time, self.end_time, self.print_verbose + ) + + # Verify print_verbose was called with error message + self.print_verbose.assert_called_once_with( + "Athina Logger Error - Bad Request, 400" + ) + + @patch("litellm.module_level_client.post") + def test_log_event_exception(self, mock_post): + """Test handling of exceptions during logging""" + # Setup mock to raise exception + mock_post.side_effect = Exception("Test exception") + + # Create test data + kwargs = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "stream": False, + } + + response_obj = MagicMock() + response_obj.model_dump.return_value = {} + + # Call the method + self.logger.log_event( + kwargs, response_obj, self.start_time, self.end_time, self.print_verbose + ) + + # Verify print_verbose was called with exception info + self.print_verbose.assert_called_once() + self.assertIn( + "Athina Logger Error - Test exception", self.print_verbose.call_args[0][0] + ) + + @patch("litellm.module_level_client.post") + def test_log_event_with_tools(self, mock_post): + """Test logging with tools/functions data""" + # Setup mock response + mock_response = MagicMock() + mock_response.status_code = 200 + mock_post.return_value = mock_response + + # Create test data with tools + kwargs = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "What's the weather?"}], + "stream": False, + "optional_params": { + "tools": [{"type": "function", "function": {"name": "get_weather"}}] + }, + } + + response_obj = MagicMock() + response_obj.model_dump.return_value = { + "id": "resp-123", + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + + # Call the method + self.logger.log_event( + kwargs, response_obj, self.start_time, self.end_time, self.print_verbose + ) + + # Verify the results + sent_data = json.loads(mock_post.call_args[1]["data"]) + self.assertEqual( + sent_data["tools"], + [{"type": "function", "function": {"name": "get_weather"}}], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/litellm/integrations/test_custom_prompt_management.py b/tests/test_litellm/integrations/test_custom_prompt_management.py similarity index 100% rename from tests/litellm/integrations/test_custom_prompt_management.py rename to tests/test_litellm/integrations/test_custom_prompt_management.py diff --git a/tests/litellm/integrations/test_deepeval.py b/tests/test_litellm/integrations/test_deepeval.py similarity index 95% rename from tests/litellm/integrations/test_deepeval.py rename to tests/test_litellm/integrations/test_deepeval.py index e62e799e301..26e79585c91 100644 --- a/tests/litellm/integrations/test_deepeval.py +++ b/tests/test_litellm/integrations/test_deepeval.py @@ -1,12 +1,12 @@ -import unittest -from unittest.mock import patch, MagicMock -from datetime import datetime, timezone -import uuid import os +import unittest +import uuid +from datetime import datetime, timezone +from unittest.mock import MagicMock, patch +from litellm.integrations.deepeval.api import Endpoints, HttpMethods from litellm.integrations.deepeval.deepeval import DeepEvalLogger -from litellm.integrations.deepeval.api import HttpMethods, Endpoints -from litellm.integrations.deepeval.types import TraceSpanApiStatus, SpanApiType +from litellm.integrations.deepeval.types import SpanApiType, TraceSpanApiStatus class TestDeepEvalLogger(unittest.TestCase): diff --git a/tests/litellm/integrations/test_langfuse.py b/tests/test_litellm/integrations/test_langfuse.py similarity index 100% rename from tests/litellm/integrations/test_langfuse.py rename to tests/test_litellm/integrations/test_langfuse.py diff --git a/tests/litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py similarity index 100% rename from tests/litellm/integrations/test_opentelemetry.py rename to tests/test_litellm/integrations/test_opentelemetry.py diff --git a/tests/litellm/integrations/test_prometheus.py b/tests/test_litellm/integrations/test_prometheus.py similarity index 100% rename from tests/litellm/integrations/test_prometheus.py rename to tests/test_litellm/integrations/test_prometheus.py diff --git a/tests/litellm/integrations/test_prometheus_services.py b/tests/test_litellm/integrations/test_prometheus_services.py similarity index 100% rename from tests/litellm/integrations/test_prometheus_services.py rename to tests/test_litellm/integrations/test_prometheus_services.py diff --git a/tests/test_litellm/litellm_core_utils/__init__.py b/tests/test_litellm/litellm_core_utils/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py similarity index 100% rename from tests/litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py rename to tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py diff --git a/tests/litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py similarity index 100% rename from tests/litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py rename to tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py diff --git a/tests/litellm/litellm_core_utils/messages_with_counts.py b/tests/test_litellm/litellm_core_utils/messages_with_counts.py similarity index 99% rename from tests/litellm/litellm_core_utils/messages_with_counts.py rename to tests/test_litellm/litellm_core_utils/messages_with_counts.py index 4595a94024c..534f8e20c2a 100644 --- a/tests/litellm/litellm_core_utils/messages_with_counts.py +++ b/tests/test_litellm/litellm_core_utils/messages_with_counts.py @@ -410,8 +410,8 @@ inner_object = { } ], "tool_choice": "none", - "count": 65, - "count-tolerate" : 67 #over by 2 + "count": 65, + "count-tolerate": 67, # over by 2 } """ namespace functions { @@ -459,7 +459,7 @@ inner_object_with_enum_only = { ], "tool_choice": "none", "count": 73, - "count-tolerate" : 74 #over by 1 + "count-tolerate": 74, # over by 1 } """ namespace functions { @@ -511,7 +511,7 @@ inner_object_with_enum = { ], "tool_choice": "none", "count": 89, - "count-tolerate" : 92, #over by 3 + "count-tolerate": 92, # over by 3 } """ namespace functions { @@ -568,8 +568,8 @@ inner_object_and_string = { } ], "tool_choice": "none", - "count": 103, - "count-tolerate" : 106, #over by 3 + "count": 103, + "count-tolerate": 106, # over by 3 } """ namespace functions { diff --git a/tests/litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py similarity index 100% rename from tests/litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py rename to tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py diff --git a/tests/litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py similarity index 97% rename from tests/litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py rename to tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index f78b933c92e..a8b9cd8a44c 100644 --- a/tests/litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -6,10 +6,11 @@ import pytest import litellm from litellm.litellm_core_utils.prompt_templates.factory import ( BAD_MESSAGE_ERROR_STR, - ollama_pt, BedrockConverseMessagesProcessor, + ollama_pt, ) + def test_ollama_pt_simple_messages(): """Test basic functionality with simple text messages""" messages = [ @@ -43,6 +44,7 @@ def test_ollama_pt_consecutive_user_messages(): assert isinstance(result, dict) assert result["prompt"] == expected_prompt + @pytest.mark.asyncio async def test_anthropic_bedrock_thinking_blocks_with_none_content(): """ @@ -55,28 +57,29 @@ async def test_anthropic_bedrock_thinking_blocks_with_none_content(): { "type": "thinking", "thinking": "This is a test thinking block", - "signature": "test-signature" + "signature": "test-signature", } ], - "reasoning_content": "This is the reasoning content" + "reasoning_content": "This is the reasoning content", } messages = [ {"role": "user", "content": "What is the capital of France?"}, - mock_assistant_message + mock_assistant_message, ] # test _bedrock_converse_messages_pt_async result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( messages=messages, model="us.anthropic.claude-3-7-sonnet-20250219-v1:0", - llm_provider="bedrock" + llm_provider="bedrock", ) - # verify the result assert len(result) == 2 - assert result[1]["content"][0]["reasoningContent"]["reasoningText"]["text"] == "This is a test thinking block" - + assert ( + result[1]["content"][0]["reasoningContent"]["reasoningText"]["text"] + == "This is a test thinking block" + ) # def test_ollama_pt_consecutive_system_messages(): diff --git a/tests/litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py similarity index 100% rename from tests/litellm/litellm_core_utils/test_core_helpers.py rename to tests/test_litellm/litellm_core_utils/test_core_helpers.py diff --git a/tests/litellm/litellm_core_utils/test_dd_tracing.py b/tests/test_litellm/litellm_core_utils/test_dd_tracing.py similarity index 100% rename from tests/litellm/litellm_core_utils/test_dd_tracing.py rename to tests/test_litellm/litellm_core_utils/test_dd_tracing.py diff --git a/tests/litellm/litellm_core_utils/test_duration_parser.py b/tests/test_litellm/litellm_core_utils/test_duration_parser.py similarity index 92% rename from tests/litellm/litellm_core_utils/test_duration_parser.py rename to tests/test_litellm/litellm_core_utils/test_duration_parser.py index 52b0f496497..b6e73dab6bc 100644 --- a/tests/litellm/litellm_core_utils/test_duration_parser.py +++ b/tests/test_litellm/litellm_core_utils/test_duration_parser.py @@ -1,30 +1,34 @@ import unittest from datetime import datetime, timezone from zoneinfo import ZoneInfo + from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time + class TestStandardizedResetTime(unittest.TestCase): def test_day_based_resets(self): """Test day-based reset durations (1d, 7d, 30d)""" # Base time: 2023-05-15 10:30:00 UTC base_time = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc) - + # Daily reset (1d) - should reset at next midnight daily_expected = datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc) daily_result = get_next_standardized_reset_time("1d", base_time, "UTC") self.assertEqual(daily_result, daily_expected) - + # Weekly reset (7d) - should reset on next Monday wednesday = datetime(2023, 5, 17, 15, 45, 0, tzinfo=timezone.utc) # A Wednesday - weekly_expected = datetime(2023, 5, 22, 0, 0, 0, tzinfo=timezone.utc) # Next Monday + weekly_expected = datetime( + 2023, 5, 22, 0, 0, 0, tzinfo=timezone.utc + ) # Next Monday weekly_result = get_next_standardized_reset_time("7d", wednesday, "UTC") self.assertEqual(weekly_result, weekly_expected) - + # Monthly reset (30d) - should reset on 1st of next month monthly_expected = datetime(2023, 6, 1, 0, 0, 0, tzinfo=timezone.utc) monthly_result = get_next_standardized_reset_time("30d", base_time, "UTC") self.assertEqual(monthly_result, monthly_expected) - + # Custom day reset (3d) - should reset after 3 days custom_day_expected = datetime(2023, 5, 18, 0, 0, 0, tzinfo=timezone.utc) custom_day_result = get_next_standardized_reset_time("3d", base_time, "UTC") @@ -34,17 +38,17 @@ class TestStandardizedResetTime(unittest.TestCase): """Test hour, minute, and second based reset durations""" # Base time: 2023-05-15 15:20:30 UTC (3:20:30 PM) base_time = datetime(2023, 5, 15, 15, 20, 30, tzinfo=timezone.utc) - + # 2-hour reset - should reset at next even hour (16:00) hour_expected = datetime(2023, 5, 15, 16, 0, 0, tzinfo=timezone.utc) hour_result = get_next_standardized_reset_time("2h", base_time, "UTC") self.assertEqual(hour_result, hour_expected) - + # 30-minute reset - should reset at next 30-minute mark (15:30) minute_expected = datetime(2023, 5, 15, 15, 30, 0, tzinfo=timezone.utc) minute_result = get_next_standardized_reset_time("30m", base_time, "UTC") self.assertEqual(minute_result, minute_expected) - + # 15-second reset - should reset at next 15-second mark (15:20:45) second_expected = datetime(2023, 5, 15, 15, 20, 45, tzinfo=timezone.utc) second_result = get_next_standardized_reset_time("15s", base_time, "UTC") @@ -54,32 +58,34 @@ class TestStandardizedResetTime(unittest.TestCase): """Test timezone handling with different regions""" # Base time: 2023-05-15 22:30:00 UTC (late in UTC day) base_time = datetime(2023, 5, 15, 22, 30, 0, tzinfo=timezone.utc) - + # Test daily reset in different timezones # US/Eastern (UTC-4): 6:30 PM, so next reset is midnight same day eastern = ZoneInfo("US/Eastern") eastern_expected = datetime(2023, 5, 16, 0, 0, 0, tzinfo=eastern) eastern_result = get_next_standardized_reset_time("1d", base_time, "US/Eastern") self.assertEqual(eastern_result, eastern_expected) - + # Asia/Kolkata (UTC+5:30): 4:00 AM next day, so next reset is midnight the day after ist = ZoneInfo("Asia/Kolkata") ist_expected = datetime(2023, 5, 17, 0, 0, 0, tzinfo=ist) ist_result = get_next_standardized_reset_time("1d", base_time, "Asia/Kolkata") self.assertEqual(ist_result, ist_expected) - + # Test hourly reset in different timezones # US/Pacific (UTC-7): 3:30 PM, so next 2h reset is 4:00 PM pacific = ZoneInfo("US/Pacific") pacific_expected = datetime(2023, 5, 15, 16, 0, 0, tzinfo=pacific) pacific_result = get_next_standardized_reset_time("2h", base_time, "US/Pacific") self.assertEqual(pacific_result, pacific_expected) - + # Test minute reset in different timezones # Europe/London (UTC+1): 11:30 PM, so next 15m reset is 11:45 PM london = ZoneInfo("Europe/London") london_expected = datetime(2023, 5, 15, 23, 45, 0, tzinfo=london) - london_result = get_next_standardized_reset_time("15m", base_time, "Europe/London") + london_result = get_next_standardized_reset_time( + "15m", base_time, "Europe/London" + ) self.assertEqual(london_result, london_expected) def test_edge_cases(self): @@ -89,25 +95,30 @@ class TestStandardizedResetTime(unittest.TestCase): hour_expected = datetime(2023, 5, 15, 16, 0, 0, tzinfo=timezone.utc) hour_result = get_next_standardized_reset_time("2h", on_hour, "UTC") self.assertEqual(hour_result, hour_expected) - + # Exactly on minute boundary on_minute = datetime(2023, 5, 15, 14, 30, 0, tzinfo=timezone.utc) minute_expected = datetime(2023, 5, 15, 15, 0, 0, tzinfo=timezone.utc) minute_result = get_next_standardized_reset_time("30m", on_minute, "UTC") self.assertEqual(minute_result, minute_expected) - + # Near day boundary near_midnight = datetime(2023, 5, 15, 23, 50, 0, tzinfo=timezone.utc) - + # 30m near midnight - should roll over to next day midnight_minute_expected = datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc) - midnight_minute_result = get_next_standardized_reset_time("30m", near_midnight, "UTC") + midnight_minute_result = get_next_standardized_reset_time( + "30m", near_midnight, "UTC" + ) self.assertEqual(midnight_minute_result, midnight_minute_expected) - + # Invalid timezone - should fall back to UTC invalid_tz_expected = datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc) - invalid_tz_result = get_next_standardized_reset_time("1d", on_hour, "NonExistentTimeZone") + invalid_tz_result = get_next_standardized_reset_time( + "1d", on_hour, "NonExistentTimeZone" + ) self.assertEqual(invalid_tz_result, invalid_tz_expected) + if __name__ == "__main__": - unittest.main() \ No newline at end of file + unittest.main() diff --git a/tests/litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py similarity index 100% rename from tests/litellm/litellm_core_utils/test_litellm_logging.py rename to tests/test_litellm/litellm_core_utils/test_litellm_logging.py diff --git a/tests/litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py similarity index 100% rename from tests/litellm/litellm_core_utils/test_realtime_streaming.py rename to tests/test_litellm/litellm_core_utils/test_realtime_streaming.py diff --git a/tests/litellm/litellm_core_utils/test_safe_json_dumps.py b/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py similarity index 100% rename from tests/litellm/litellm_core_utils/test_safe_json_dumps.py rename to tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py diff --git a/tests/litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py similarity index 100% rename from tests/litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py rename to tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py index d1a452b80a7..dfc3f01d112 100644 --- a/tests/litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -15,9 +15,9 @@ from litellm.types.utils import ( Delta, Function, ModelResponseStream, + PromptTokensDetails, StreamingChoices, Usage, - PromptTokensDetails, ) diff --git a/tests/litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py similarity index 97% rename from tests/litellm/litellm_core_utils/test_streaming_handler.py rename to tests/test_litellm/litellm_core_utils/test_streaming_handler.py index c38709735f0..8027db94428 100644 --- a/tests/litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -612,29 +612,45 @@ def test_streaming_handler_with_stop_chunk( assert returned_chunk is None -def test_set_response_id_propagation_empty_to_valid(initialized_custom_stream_wrapper: CustomStreamWrapper): +def test_set_response_id_propagation_empty_to_valid( + initialized_custom_stream_wrapper: CustomStreamWrapper, +): """Test that response_id is properly set when first chunk has empty ID and second chunk has valid ID""" model_response1 = ModelResponseStream(id="", created=1742056047, model=None) - model_response1 = initialized_custom_stream_wrapper.set_model_id(model_response1.id, model_response1) + model_response1 = initialized_custom_stream_wrapper.set_model_id( + model_response1.id, model_response1 + ) assert model_response1.id == "" - model_response2 = ModelResponseStream(id="valid-id-123", created=1742056048, model=None) - model_response2 = initialized_custom_stream_wrapper.set_model_id("valid-id-123", model_response2) + model_response2 = ModelResponseStream( + id="valid-id-123", created=1742056048, model=None + ) + model_response2 = initialized_custom_stream_wrapper.set_model_id( + "valid-id-123", model_response2 + ) assert model_response2.id == "valid-id-123" assert initialized_custom_stream_wrapper.response_id == "valid-id-123" -def test_set_response_id_propagation_valid_to_invalid(initialized_custom_stream_wrapper: CustomStreamWrapper): +def test_set_response_id_propagation_valid_to_invalid( + initialized_custom_stream_wrapper: CustomStreamWrapper, +): """Test that response_id is maintained when first chunk has valid ID and second chunk has invalid ID""" - model_response1 = ModelResponseStream(id="first-valid-id", created=1742056049, model=None) - model_response1 = initialized_custom_stream_wrapper.set_model_id("first-valid-id", model_response1) + model_response1 = ModelResponseStream( + id="first-valid-id", created=1742056049, model=None + ) + model_response1 = initialized_custom_stream_wrapper.set_model_id( + "first-valid-id", model_response1 + ) assert model_response1.id == "first-valid-id" assert initialized_custom_stream_wrapper.response_id == "first-valid-id" model_response2 = ModelResponseStream(id="", created=1742056050, model=None) - model_response2 = initialized_custom_stream_wrapper.set_model_id("", model_response2) + model_response2 = initialized_custom_stream_wrapper.set_model_id( + "", model_response2 + ) assert model_response2.id == "first-valid-id" assert initialized_custom_stream_wrapper.response_id == "first-valid-id" diff --git a/tests/litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py similarity index 99% rename from tests/litellm/litellm_core_utils/test_token_counter.py rename to tests/test_litellm/litellm_core_utils/test_token_counter.py index 07c1367e5f8..b2d3bcd1305 100644 --- a/tests/litellm/litellm_core_utils/test_token_counter.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter.py @@ -13,17 +13,16 @@ sys.path.insert( ) # Adds the parent directory to the system path from unittest.mock import AsyncMock, MagicMock, patch -from messages_with_counts import ( - MESSAGES_TEXT, - MESSAGES_WITH_IMAGES, - MESSAGES_WITH_TOOLS, -) - import litellm from litellm import create_pretrained_tokenizer, decode, encode, get_modified_max_tokens from litellm import token_counter as token_counter_old from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new from tests.large_text import text +from tests.test_litellm.litellm_core_utils.messages_with_counts import ( + MESSAGES_TEXT, + MESSAGES_WITH_IMAGES, + MESSAGES_WITH_TOOLS, +) def token_counter_both_assert_same(**args): diff --git a/tests/litellm/litellm_core_utils/test_token_counter_tool.py b/tests/test_litellm/litellm_core_utils/test_token_counter_tool.py similarity index 92% rename from tests/litellm/litellm_core_utils/test_token_counter_tool.py rename to tests/test_litellm/litellm_core_utils/test_token_counter_tool.py index 36cab0bc50c..e8836bab2b9 100644 --- a/tests/litellm/litellm_core_utils/test_token_counter_tool.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter_tool.py @@ -9,10 +9,9 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path -#Use the same token_counter as the main test. -from test_token_counter import token_counter - -from test_token_counter_tool_data import * +# Use the same token_counter as the main test. +from tests.test_litellm.litellm_core_utils.test_token_counter import token_counter +from tests.test_litellm.litellm_core_utils.test_token_counter_tool_data import * @pytest.mark.parametrize( diff --git a/tests/litellm/litellm_core_utils/test_token_counter_tool_data.py b/tests/test_litellm/litellm_core_utils/test_token_counter_tool_data.py similarity index 86% rename from tests/litellm/litellm_core_utils/test_token_counter_tool_data.py rename to tests/test_litellm/litellm_core_utils/test_token_counter_tool_data.py index bf030b673be..c58cff6a3f2 100644 --- a/tests/litellm/litellm_core_utils/test_token_counter_tool_data.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter_tool_data.py @@ -27,19 +27,19 @@ CONTENT_AND_TOOL_CALL = [ ] _OPENHANDS_SYSTEM_MESSAGE = { - "content": [ - { - "type": "text", - "text": "You are OpenHands agent, a helpful AI assistant that can " - "interact with a computer to solve tasks.\n\n\nYour primary " - "role is to assist users by executing commands, modifying code, and " - "solving technical problems effectively. You should be thorough, " - "methodical, and prioritize quality over speed.\n* If the user asks a " - "question, like 'why is X happening', don’t try to fix the problem. ", - } - ], - "role": "system", - } + "content": [ + { + "type": "text", + "text": "You are OpenHands agent, a helpful AI assistant that can " + "interact with a computer to solve tasks.\n\n\nYour primary " + "role is to assist users by executing commands, modifying code, and " + "solving technical problems effectively. You should be thorough, " + "methodical, and prioritize quality over speed.\n* If the user asks a " + "question, like 'why is X happening', don’t try to fix the problem. ", + } + ], + "role": "system", +} SYSTEM_LONG = [ _OPENHANDS_SYSTEM_MESSAGE, diff --git a/tests/litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py similarity index 100% rename from tests/litellm/llms/anthropic/chat/test_anthropic_chat_handler.py rename to tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py diff --git a/tests/litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py similarity index 100% rename from tests/litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py rename to tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py diff --git a/tests/litellm/llms/azure/image_generation/test_azure_image_generation_init.py b/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py similarity index 100% rename from tests/litellm/llms/azure/image_generation/test_azure_image_generation_init.py rename to tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py diff --git a/tests/litellm/llms/azure/test_azure_common_utils.py b/tests/test_litellm/llms/azure/test_azure_common_utils.py similarity index 100% rename from tests/litellm/llms/azure/test_azure_common_utils.py rename to tests/test_litellm/llms/azure/test_azure_common_utils.py diff --git a/tests/litellm/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py similarity index 68% rename from tests/litellm/llms/azure_ai/chat/test_azure_ai_transformation.py rename to tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py index 6a42b51fe91..ffb485c3680 100644 --- a/tests/litellm/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -19,13 +19,15 @@ async def test_get_openai_compatible_provider_info(): """ config = AzureAIStudioConfig() - api_base, dynamic_api_key, custom_llm_provider = ( - config._get_openai_compatible_provider_info( - model="azure_ai/gpt-4o-mini", - api_base="https://my-base", - api_key="my-key", - custom_llm_provider="azure_ai", - ) + ( + api_base, + dynamic_api_key, + custom_llm_provider, + ) = config._get_openai_compatible_provider_info( + model="azure_ai/gpt-4o-mini", + api_base="https://my-base", + api_key="my-key", + custom_llm_provider="azure_ai", ) assert custom_llm_provider == "azure" diff --git a/tests/litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py similarity index 100% rename from tests/litellm/llms/bedrock/chat/test_converse_transformation.py rename to tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py diff --git a/tests/litellm/llms/bedrock/chat/test_invoke_handler.py b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py similarity index 100% rename from tests/litellm/llms/bedrock/chat/test_invoke_handler.py rename to tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py diff --git a/tests/litellm/llms/bedrock/chat/test_mistral_config.py b/tests/test_litellm/llms/bedrock/chat/test_mistral_config.py similarity index 85% rename from tests/litellm/llms/bedrock/chat/test_mistral_config.py rename to tests/test_litellm/llms/bedrock/chat/test_mistral_config.py index 02234401915..42261cb3bb1 100644 --- a/tests/litellm/llms/bedrock/chat/test_mistral_config.py +++ b/tests/test_litellm/llms/bedrock/chat/test_mistral_config.py @@ -1,6 +1,6 @@ - - -from litellm.llms.bedrock.chat.invoke_transformations.amazon_mistral_transformation import AmazonMistralConfig +from litellm.llms.bedrock.chat.invoke_transformations.amazon_mistral_transformation import ( + AmazonMistralConfig, +) from litellm.types.utils import ModelResponse @@ -10,7 +10,9 @@ def test_mistral_get_outputText(): model_response.choices[0].finish_reason = "None" # Models like pixtral will return a completion with the openai format. - mock_json_with_choices = {"choices": [{"message": {"content": "Hello!"}, "finish_reason": "stop"}]} + mock_json_with_choices = { + "choices": [{"message": {"content": "Hello!"}, "finish_reason": "stop"}] + } outputText = AmazonMistralConfig.get_outputText( completion_response=mock_json_with_choices, model_response=model_response diff --git a/tests/litellm/llms/bedrock/image/test_amazon_stability3_transformation.py b/tests/test_litellm/llms/bedrock/image/test_amazon_stability3_transformation.py similarity index 100% rename from tests/litellm/llms/bedrock/image/test_amazon_stability3_transformation.py rename to tests/test_litellm/llms/bedrock/image/test_amazon_stability3_transformation.py diff --git a/tests/litellm/llms/bedrock/rerank/transformation.py b/tests/test_litellm/llms/bedrock/rerank/transformation.py similarity index 100% rename from tests/litellm/llms/bedrock/rerank/transformation.py rename to tests/test_litellm/llms/bedrock/rerank/transformation.py diff --git a/tests/litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py similarity index 100% rename from tests/litellm/llms/bedrock/test_base_aws_llm.py rename to tests/test_litellm/llms/bedrock/test_base_aws_llm.py diff --git a/tests/litellm/llms/bedrock/test_bedrock_common_utils.py b/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py similarity index 100% rename from tests/litellm/llms/bedrock/test_bedrock_common_utils.py rename to tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py diff --git a/tests/litellm/llms/chat/test_converse_handler.py b/tests/test_litellm/llms/chat/test_converse_handler.py similarity index 100% rename from tests/litellm/llms/chat/test_converse_handler.py rename to tests/test_litellm/llms/chat/test_converse_handler.py diff --git a/tests/litellm/llms/cohere/chat/test_transformation.py b/tests/test_litellm/llms/cohere/chat/test_transformation.py similarity index 85% rename from tests/litellm/llms/cohere/chat/test_transformation.py rename to tests/test_litellm/llms/cohere/chat/test_transformation.py index 30079abd434..4fe8f8a88a9 100644 --- a/tests/litellm/llms/cohere/chat/test_transformation.py +++ b/tests/test_litellm/llms/cohere/chat/test_transformation.py @@ -2,7 +2,6 @@ import os import sys from unittest.mock import MagicMock - sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path @@ -18,7 +17,11 @@ class TestCohereTransform: def test_map_cohere_params(self): """Test that parameters are correctly mapped""" - test_params = {"temperature": 0.7, "max_tokens": 200, "max_completion_tokens": 256} + test_params = { + "temperature": 0.7, + "max_tokens": 200, + "max_completion_tokens": 256, + } result = self.config.map_openai_params( non_default_params=test_params, @@ -32,7 +35,10 @@ class TestCohereTransform: def test_cohere_max_tokens_backward_compat(self): """Test that parameters are correctly mapped""" - test_params = {"temperature": 0.7, "max_tokens": 200,} + test_params = { + "temperature": 0.7, + "max_tokens": 200, + } result = self.config.map_openai_params( non_default_params=test_params, diff --git a/tests/litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py similarity index 100% rename from tests/litellm/llms/custom_httpx/test_http_handler.py rename to tests/test_litellm/llms/custom_httpx/test_http_handler.py diff --git a/tests/litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py similarity index 100% rename from tests/litellm/llms/custom_httpx/test_llm_http_handler.py rename to tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py diff --git a/tests/litellm/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py similarity index 100% rename from tests/litellm/llms/databricks/chat/test_databricks_chat_transformation.py rename to tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py diff --git a/tests/litellm/llms/databricks/test_databricks_common_utils.py b/tests/test_litellm/llms/databricks/test_databricks_common_utils.py similarity index 100% rename from tests/litellm/llms/databricks/test_databricks_common_utils.py rename to tests/test_litellm/llms/databricks/test_databricks_common_utils.py diff --git a/tests/litellm/llms/deepgram/audio_transcription/test_deepseek_audio_transcription_transformation.py b/tests/test_litellm/llms/deepgram/audio_transcription/test_deepseek_audio_transcription_transformation.py similarity index 100% rename from tests/litellm/llms/deepgram/audio_transcription/test_deepseek_audio_transcription_transformation.py rename to tests/test_litellm/llms/deepgram/audio_transcription/test_deepseek_audio_transcription_transformation.py diff --git a/tests/litellm/llms/featherless_ai/chat/test_featherless_chat_transformation.py b/tests/test_litellm/llms/featherless_ai/chat/test_featherless_chat_transformation.py similarity index 100% rename from tests/litellm/llms/featherless_ai/chat/test_featherless_chat_transformation.py rename to tests/test_litellm/llms/featherless_ai/chat/test_featherless_chat_transformation.py diff --git a/tests/litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py similarity index 100% rename from tests/litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py rename to tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py diff --git a/tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py b/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py similarity index 100% rename from tests/litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py rename to tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py diff --git a/tests/litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py b/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py similarity index 100% rename from tests/litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py rename to tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py diff --git a/tests/litellm/llms/llamafile/chat/test_llamafile_chat_transformation.py b/tests/test_litellm/llms/llamafile/chat/test_llamafile_chat_transformation.py similarity index 52% rename from tests/litellm/llms/llamafile/chat/test_llamafile_chat_transformation.py rename to tests/test_litellm/llms/llamafile/chat/test_llamafile_chat_transformation.py index 3dfd0d66286..650f86f4e49 100644 --- a/tests/litellm/llms/llamafile/chat/test_llamafile_chat_transformation.py +++ b/tests/test_litellm/llms/llamafile/chat/test_llamafile_chat_transformation.py @@ -1,6 +1,6 @@ from typing import Optional - from unittest.mock import patch + import pytest import litellm @@ -13,12 +13,26 @@ from litellm.llms.llamafile.chat.transformation import LlamafileChatConfig ("user-provided-key", "secret-key", "user-provided-key", False), (None, "secret-key", "secret-key", True), (None, None, "fake-api-key", True), - ("", "secret-key", "secret-key", True), # Empty string should fall back to secret - ("", None, "fake-api-key", True), # Empty string with no secret should use the fake key - ] + ( + "", + "secret-key", + "secret-key", + True, + ), # Empty string should fall back to secret + ( + "", + None, + "fake-api-key", + True, + ), # Empty string with no secret should use the fake key + ], ) -def test_resolve_api_key(input_api_key, api_key_from_secret_manager, expected_api_key, secret_manager_called): - with patch("litellm.llms.llamafile.chat.transformation.get_secret_str") as mock_get_secret: +def test_resolve_api_key( + input_api_key, api_key_from_secret_manager, expected_api_key, secret_manager_called +): + with patch( + "litellm.llms.llamafile.chat.transformation.get_secret_str" + ) as mock_get_secret: mock_get_secret.return_value = api_key_from_secret_manager result = LlamafileChatConfig._resolve_api_key(input_api_key) @@ -34,14 +48,36 @@ def test_resolve_api_key(input_api_key, api_key_from_secret_manager, expected_ap @pytest.mark.parametrize( "input_api_base, api_base_from_secret_manager, expected_api_base, secret_manager_called", [ - ("https://user-api.example.com", "https://secret-api.example.com", "https://user-api.example.com", False), - (None, "https://secret-api.example.com", "https://secret-api.example.com", True), + ( + "https://user-api.example.com", + "https://secret-api.example.com", + "https://user-api.example.com", + False, + ), + ( + None, + "https://secret-api.example.com", + "https://secret-api.example.com", + True, + ), (None, None, "http://127.0.0.1:8080/v1", True), - ("", "https://secret-api.example.com", "https://secret-api.example.com", True), # Empty string should fall back - ] + ( + "", + "https://secret-api.example.com", + "https://secret-api.example.com", + True, + ), # Empty string should fall back + ], ) -def test_resolve_api_base(input_api_base, api_base_from_secret_manager, expected_api_base, secret_manager_called): - with patch("litellm.llms.llamafile.chat.transformation.get_secret_str") as mock_get_secret: +def test_resolve_api_base( + input_api_base, + api_base_from_secret_manager, + expected_api_base, + secret_manager_called, +): + with patch( + "litellm.llms.llamafile.chat.transformation.get_secret_str" + ) as mock_get_secret: mock_get_secret.return_value = api_base_from_secret_manager result = LlamafileChatConfig._resolve_api_base(input_api_base) @@ -58,31 +94,73 @@ def test_resolve_api_base(input_api_base, api_base_from_secret_manager, expected "api_base, api_key, secret_base, secret_key, expected_base, expected_key", [ # User-provided values - ("https://user-api.example.com", "user-key", "https://secret-api.example.com", "secret-key", "https://user-api.example.com", "user-key"), + ( + "https://user-api.example.com", + "user-key", + "https://secret-api.example.com", + "secret-key", + "https://user-api.example.com", + "user-key", + ), # Fallback to secrets - (None, None, "https://secret-api.example.com", "secret-key", "https://secret-api.example.com", "secret-key"), + ( + None, + None, + "https://secret-api.example.com", + "secret-key", + "https://secret-api.example.com", + "secret-key", + ), # Nothing provided, use defaults (None, None, None, None, "http://127.0.0.1:8080/v1", "fake-api-key"), # Mixed scenarios - ("https://user-api.example.com", None, None, "secret-key", "https://user-api.example.com", "secret-key"), - (None, "user-key", "https://secret-api.example.com", None, "https://secret-api.example.com", "user-key"), - ] + ( + "https://user-api.example.com", + None, + None, + "secret-key", + "https://user-api.example.com", + "secret-key", + ), + ( + None, + "user-key", + "https://secret-api.example.com", + None, + "https://secret-api.example.com", + "user-key", + ), + ], ) -def test_get_openai_compatible_provider_info(api_base, api_key, secret_base, secret_key, expected_base, expected_key): +def test_get_openai_compatible_provider_info( + api_base, api_key, secret_base, secret_key, expected_base, expected_key +): config = LlamafileChatConfig() def fake_get_secret(key: str) -> Optional[str]: - return { - "LLAMAFILE_API_BASE": secret_base, - "LLAMAFILE_API_KEY": secret_key - }.get(key) + return {"LLAMAFILE_API_BASE": secret_base, "LLAMAFILE_API_KEY": secret_key}.get( + key + ) - patch_secret = patch("litellm.llms.llamafile.chat.transformation.get_secret_str", side_effect=fake_get_secret) - patch_base = patch.object(LlamafileChatConfig, "_resolve_api_base", wraps=LlamafileChatConfig._resolve_api_base) - patch_key = patch.object(LlamafileChatConfig, "_resolve_api_key", wraps=LlamafileChatConfig._resolve_api_key) + patch_secret = patch( + "litellm.llms.llamafile.chat.transformation.get_secret_str", + side_effect=fake_get_secret, + ) + patch_base = patch.object( + LlamafileChatConfig, + "_resolve_api_base", + wraps=LlamafileChatConfig._resolve_api_base, + ) + patch_key = patch.object( + LlamafileChatConfig, + "_resolve_api_key", + wraps=LlamafileChatConfig._resolve_api_key, + ) with patch_secret as mock_secret, patch_base as mock_base, patch_key as mock_key: - result_base, result_key = config._get_openai_compatible_provider_info(api_base, api_key) + result_base, result_key = config._get_openai_compatible_provider_info( + api_base, api_key + ) assert result_base == expected_base assert result_key == expected_key @@ -100,8 +178,12 @@ def test_get_openai_compatible_provider_info(api_base, api_key, secret_base, sec def test_completion_with_custom_llamafile_model(): - with patch("litellm.main.openai_chat_completions.completion") as mock_llamafile_completion_func: - mock_llamafile_completion_func.return_value = {} # Return an empty dictionary for the mocked response + with patch( + "litellm.main.openai_chat_completions.completion" + ) as mock_llamafile_completion_func: + mock_llamafile_completion_func.return_value = ( + {} + ) # Return an empty dictionary for the mocked response provider = "llamafile" model_name = "my-custom-test-model" diff --git a/tests/litellm/llms/lm_studio/test_lm_studio_chat_transformation.py b/tests/test_litellm/llms/lm_studio/test_lm_studio_chat_transformation.py similarity index 88% rename from tests/litellm/llms/lm_studio/test_lm_studio_chat_transformation.py rename to tests/test_litellm/llms/lm_studio/test_lm_studio_chat_transformation.py index 718142ad175..1f09b1e5406 100644 --- a/tests/litellm/llms/lm_studio/test_lm_studio_chat_transformation.py +++ b/tests/test_litellm/llms/lm_studio/test_lm_studio_chat_transformation.py @@ -3,7 +3,9 @@ import sys from pydantic import BaseModel -sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))) +sys.path.insert( + 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../..")) +) from litellm.llms.lm_studio.chat.transformation import LMStudioChatConfig from litellm.utils import get_optional_params @@ -44,7 +46,9 @@ class TestLMStudioChatConfigResponseFormat: custom_llm_provider="lm_studio", ) - mapped = config.map_openai_params(non_default_params, {}, "lm_studio/test-model", False) + mapped = config.map_openai_params( + non_default_params, {}, "lm_studio/test-model", False + ) mapped_schema = mapped["response_format"]["json_schema"]["schema"] assert mapped_schema["properties"] == schema["properties"] opt_schema = optional_params["response_format"]["json_schema"]["schema"] diff --git a/tests/litellm/llms/meta_llama/test_meta_llama_chat_transformation.py b/tests/test_litellm/llms/meta_llama/test_meta_llama_chat_transformation.py similarity index 100% rename from tests/litellm/llms/meta_llama/test_meta_llama_chat_transformation.py rename to tests/test_litellm/llms/meta_llama/test_meta_llama_chat_transformation.py diff --git a/tests/litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py similarity index 100% rename from tests/litellm/llms/mistral/test_mistral_chat_transformation.py rename to tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py diff --git a/tests/litellm/llms/novita/chat/test_novita_chat_transformation.py b/tests/test_litellm/llms/novita/chat/test_novita_chat_transformation.py similarity index 100% rename from tests/litellm/llms/novita/chat/test_novita_chat_transformation.py rename to tests/test_litellm/llms/novita/chat/test_novita_chat_transformation.py diff --git a/tests/litellm/llms/nscale/chat/test_nscale_chat_transformation.py b/tests/test_litellm/llms/nscale/chat/test_nscale_chat_transformation.py similarity index 100% rename from tests/litellm/llms/nscale/chat/test_nscale_chat_transformation.py rename to tests/test_litellm/llms/nscale/chat/test_nscale_chat_transformation.py diff --git a/tests/litellm/llms/ollama/test_ollama_chat_transformation.py b/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py similarity index 100% rename from tests/litellm/llms/ollama/test_ollama_chat_transformation.py rename to tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py diff --git a/tests/litellm/llms/ollama/test_ollama_completion_transformation.py b/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py similarity index 81% rename from tests/litellm/llms/ollama/test_ollama_completion_transformation.py rename to tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py index e724bc35f83..f0b5c00d017 100644 --- a/tests/litellm/llms/ollama/test_ollama_completion_transformation.py +++ b/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py @@ -1,45 +1,42 @@ +import json import os import sys -import json import uuid -import pytest from unittest.mock import MagicMock, patch +import pytest sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path -from litellm.llms.ollama.completion.transformation import ( - OllamaConfig, -) -from litellm.types.utils import ModelResponse -from litellm.types.utils import Message +from litellm.llms.ollama.completion.transformation import OllamaConfig +from litellm.types.utils import Message, ModelResponse class TestOllamaConfig: def test_transform_response_standard(self): # Initialize config config = OllamaConfig() - + # Create mock response raw_response = MagicMock() raw_response.json.return_value = { "response": "Hello, I am an AI assistant", "prompt_eval_count": 10, - "eval_count": 5 + "eval_count": 5, } - + # Create properly structured model response object model_response = ModelResponse( id="test_id", choices=[{"message": Message(content="")}], ) - + # Create mock encoding mock_encoding = MagicMock() mock_encoding.encode.return_value = [1, 2, 3] # Return dummy token IDs - + # Transform response result = config.transform_response( model="llama2", @@ -52,7 +49,7 @@ class TestOllamaConfig: litellm_params={}, encoding=mock_encoding, ) - + # Verify response assert result.choices[0]["message"].content == "Hello, I am an AI assistant" assert result.choices[0]["finish_reason"] == "stop" @@ -67,29 +64,28 @@ class TestOllamaConfig: def test_transform_response_json_function_call(self, mock_uuid4): # Setup mock UUID mock_uuid4.return_value = "test-uuid" - + # Initialize config config = OllamaConfig() - + # Create mock response with JSON function call format raw_response = MagicMock() raw_response.json.return_value = { - "response": json.dumps({ - "name": "get_weather", - "arguments": {"location": "San Francisco"} - }) + "response": json.dumps( + {"name": "get_weather", "arguments": {"location": "San Francisco"}} + ) } - + # Create properly structured model response object model_response = ModelResponse( id="test_id", choices=[{"message": Message(content="")}], ) - + # Create mock encoding mock_encoding = MagicMock() mock_encoding.encode.return_value = [1, 2, 3] # Return dummy token IDs - + # Transform response result = config.transform_response( model="llama2", @@ -102,39 +98,43 @@ class TestOllamaConfig: litellm_params={}, encoding=mock_encoding, ) - + # Verify result has tool_calls assert result.choices[0]["message"].content is None assert result.choices[0]["finish_reason"] == "tool_calls" assert len(result.choices[0]["message"].tool_calls) == 1 assert result.choices[0]["message"].tool_calls[0]["id"].startswith("call_") - assert result.choices[0]["message"].tool_calls[0]["function"]["name"] == "get_weather" - assert json.loads(result.choices[0]["message"].tool_calls[0]["function"]["arguments"]) == {"location": "San Francisco"} + assert ( + result.choices[0]["message"].tool_calls[0]["function"]["name"] + == "get_weather" + ) + assert json.loads( + result.choices[0]["message"].tool_calls[0]["function"]["arguments"] + ) == {"location": "San Francisco"} # No usage assertions here as we don't need to test them in every case def test_transform_response_regular_json(self): # Initialize config config = OllamaConfig() - + # Create mock response with regular JSON (not function call) raw_response = MagicMock() raw_response.json.return_value = { - "response": json.dumps({ - "result": "success", - "data": {"temperature": 72, "unit": "F"} - }) + "response": json.dumps( + {"result": "success", "data": {"temperature": 72, "unit": "F"}} + ) } - + # Create properly structured model response object model_response = ModelResponse( id="test_id", choices=[{"message": Message(content="")}], ) - + # Create mock encoding mock_encoding = MagicMock() mock_encoding.encode.return_value = [1, 2, 3] # Return dummy token IDs - + # Transform response result = config.transform_response( model="llama2", @@ -147,12 +147,11 @@ class TestOllamaConfig: litellm_params={}, encoding=mock_encoding, ) - + # Verify result has JSON content - expected_content = json.dumps({ - "result": "success", - "data": {"temperature": 72, "unit": "F"} - }) + expected_content = json.dumps( + {"result": "success", "data": {"temperature": 72, "unit": "F"}} + ) assert result.choices[0]["message"].content == expected_content assert result.choices[0]["finish_reason"] == "stop" - # No usage assertions here as we don't need to test them in every case \ No newline at end of file + # No usage assertions here as we don't need to test them in every case diff --git a/tests/litellm/llms/ollama/test_ollama_model_info.py b/tests/test_litellm/llms/ollama/test_ollama_model_info.py similarity index 71% rename from tests/litellm/llms/ollama/test_ollama_model_info.py rename to tests/test_litellm/llms/ollama/test_ollama_model_info.py index f5fea572ac9..adb079763d9 100644 --- a/tests/litellm/llms/ollama/test_ollama_model_info.py +++ b/tests/test_litellm/llms/ollama/test_ollama_model_info.py @@ -1,10 +1,10 @@ +import json import os import sys -import json import uuid -import pytest from unittest.mock import MagicMock, patch +import pytest sys.path.insert( 0, os.path.abspath("../../../../..") @@ -14,13 +14,15 @@ sys.path.insert( Unit tests for OllamaModelInfo.get_models functionality. """ # Ensure a dummy httpx module is available for import in tests -import sys, types +import sys +import types + # Provide a dummy httpx module for import in get_models -if 'httpx' not in sys.modules: +if "httpx" not in sys.modules: # Create a minimal module with HTTPStatusError - httpx_mod = types.ModuleType('httpx') + httpx_mod = types.ModuleType("httpx") httpx_mod.HTTPStatusError = Exception - sys.modules['httpx'] = httpx_mod + sys.modules["httpx"] = httpx_mod import httpx @@ -31,6 +33,7 @@ class DummyResponse: """ A dummy response object to simulate httpx responses. """ + def __init__(self, json_data, status_code=200): self._json = json_data self.status_code = status_code @@ -38,7 +41,9 @@ class DummyResponse: def raise_for_status(self): if self.status_code >= 400: # Simulate an HTTP status error - raise httpx.HTTPStatusError("Error status code", request=None, response=None) + raise httpx.HTTPStatusError( + "Error status code", request=None, response=None + ) def json(self): return self._json @@ -51,25 +56,26 @@ class TestOllamaModelInfo: get_models should extract and return sorted unique model names. """ calls = [] - sample = {'models': [ - {'name': 'zeta'}, - {'model': 'alpha'}, - {'name': 123}, # non-str should be ignored - 'invalid', # non-dict should be ignored - ]} + sample = { + "models": [ + {"name": "zeta"}, + {"model": "alpha"}, + {"name": 123}, # non-str should be ignored + "invalid", # non-dict should be ignored + ] + } def mock_get(url): calls.append(url) return DummyResponse(sample, status_code=200) - monkeypatch.setattr(httpx, 'get', mock_get) + monkeypatch.setattr(httpx, "get", mock_get) info = OllamaModelInfo() models = info.get_models() # Only 'alpha' and 'zeta' should be returned, sorted alphabetically - assert models == ['alpha', 'zeta'] + assert models == ["alpha", "zeta"] # Ensure correct endpoint was called - assert calls and calls[0].endswith('/api/tags') - + assert calls and calls[0].endswith("/api/tags") def test_get_models_from_list_response(self, monkeypatch): """ @@ -77,30 +83,30 @@ class TestOllamaModelInfo: get_models should extract and return sorted unique model names. """ sample = [ - {'name': 'm1'}, - {'model': 'm2'}, - {}, # no name/model key should be ignored + {"name": "m1"}, + {"model": "m2"}, + {}, # no name/model key should be ignored ] def mock_get(url): return DummyResponse(sample, status_code=200) - monkeypatch.setattr(httpx, 'get', mock_get) + monkeypatch.setattr(httpx, "get", mock_get) info = OllamaModelInfo() models = info.get_models() - assert models == ['m1', 'm2'] - + assert models == ["m1", "m2"] def test_get_models_fallback_on_error(self, monkeypatch): """ If the httpx.get call raises an exception, get_models should fall back to the static models_by_provider list prefixed by 'ollama/'. """ + def mock_get(url): raise Exception("connection failure") - monkeypatch.setattr(httpx, 'get', mock_get) + monkeypatch.setattr(httpx, "get", mock_get) info = OllamaModelInfo() models = info.get_models() # Default static ollama_models is ['llama2'], so expect ['ollama/llama2'] - assert models == ['ollama/llama2'] \ No newline at end of file + assert models == ["ollama/llama2"] diff --git a/tests/litellm/llms/openai/responses/test_openai_responses_transformation.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py similarity index 100% rename from tests/litellm/llms/openai/responses/test_openai_responses_transformation.py rename to tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py diff --git a/tests/litellm/llms/openai/test_o_series_transformation.py b/tests/test_litellm/llms/openai/test_o_series_transformation.py similarity index 86% rename from tests/litellm/llms/openai/test_o_series_transformation.py rename to tests/test_litellm/llms/openai/test_o_series_transformation.py index 13462135195..161509a6d11 100644 --- a/tests/litellm/llms/openai/test_o_series_transformation.py +++ b/tests/test_litellm/llms/openai/test_o_series_transformation.py @@ -1,4 +1,5 @@ import pytest + from litellm.llms.openai.chat.o_series_transformation import OpenAIOSeriesConfig @@ -11,24 +12,20 @@ from litellm.llms.openai.chat.o_series_transformation import OpenAIOSeriesConfig ("o4-mini", True), ("o1-preview", True), ("o3-mini", True), - # Valid O-series models with provider prefix ("openai/o1", True), ("openai/o3", True), ("openai/o4-mini", True), ("openai/o1-preview", True), ("openai/o3-mini", True), - # Non-O-series models ("gpt-4", False), ("gpt-3.5-turbo", False), ("claude-3-opus", False), - # Non-O-series models with provider prefix ("openai/gpt-4", False), ("openai/gpt-3.5-turbo", False), ("anthropic/claude-3-opus", False), - # Edge cases ("o", False), # Too short ("o5", False), # Not a valid O-series model @@ -39,11 +36,12 @@ from litellm.llms.openai.chat.o_series_transformation import OpenAIOSeriesConfig def test_is_model_o_series_model(model_name: str, expected: bool): """ Test that is_model_o_series_model correctly identifies O-series models. - + Args: model_name: The model name to test expected: The expected result (True if it should be identified as an O-series model) """ config = OpenAIOSeriesConfig() - assert config.is_model_o_series_model(model_name) == expected, \ - f"Expected {model_name} to be {'an O-series model' if expected else 'not an O-series model'}" + assert ( + config.is_model_o_series_model(model_name) == expected + ), f"Expected {model_name} to be {'an O-series model' if expected else 'not an O-series model'}" diff --git a/tests/litellm/llms/openai/test_openai_common_utils.py b/tests/test_litellm/llms/openai/test_openai_common_utils.py similarity index 99% rename from tests/litellm/llms/openai/test_openai_common_utils.py rename to tests/test_litellm/llms/openai/test_openai_common_utils.py index a343fcf25c5..469005d103f 100644 --- a/tests/litellm/llms/openai/test_openai_common_utils.py +++ b/tests/test_litellm/llms/openai/test_openai_common_utils.py @@ -98,7 +98,6 @@ async def test_openai_client_reuse(function_name, is_async, args): ) as mock_set_cache, patch.object( BaseOpenAILLM, "get_cached_openai_client" ) as mock_get_cache: - # Setup the mock to return None first time (cache miss) then a client for subsequent calls mock_client = MagicMock() mock_get_cache.side_effect = [None] + [ diff --git a/tests/litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py b/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py similarity index 95% rename from tests/litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py rename to tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py index ecdcadfb735..7272f8c8182 100644 --- a/tests/litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py +++ b/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py @@ -3,15 +3,14 @@ import sys import pytest - sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path from litellm.llms.openrouter.chat.transformation import ( OpenRouterChatCompletionStreamingHandler, - OpenRouterException, OpenrouterConfig, + OpenRouterException, ) @@ -26,11 +25,7 @@ class TestOpenRouterChatCompletionStreamingHandler: "id": "test_id", "created": 1234567890, "model": "test_model", - "usage": { - "prompt_tokens": 10, - "completion_tokens": 20, - "total_tokens": 30 - }, + "usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, "choices": [ {"delta": {"content": "test content", "reasoning": "test reasoning"}} ], @@ -89,7 +84,6 @@ class TestOpenRouterChatCompletionStreamingHandler: def test_openrouter_extra_body_transformation(): - transformed_request = OpenrouterConfig().transform_request( model="openrouter/deepseek/deepseek-chat", messages=[{"role": "user", "content": "Hello, world!"}], diff --git a/tests/litellm/llms/sagemaker/test_sagemaker_common_utils.py b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py similarity index 95% rename from tests/litellm/llms/sagemaker/test_sagemaker_common_utils.py rename to tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py index 0d65a25b2e7..70a8d86cb1b 100644 --- a/tests/litellm/llms/sagemaker/test_sagemaker_common_utils.py +++ b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py @@ -10,6 +10,7 @@ sys.path.insert(0, os.path.abspath("../../../../..")) from litellm.llms.sagemaker.common_utils import AWSEventStreamDecoder from litellm.llms.sagemaker.completion.transformation import SagemakerConfig + @pytest.mark.asyncio async def test_aiter_bytes_unicode_decode_error(): """ @@ -96,6 +97,7 @@ async def test_aiter_bytes_valid_chunk_followed_by_unicode_error(): assert len(chunks) == 1 assert chunks[0]["text"] == "hello" # Verify the content of the valid chunk + class TestSagemakerTransform: def setup_method(self): self.config = SagemakerConfig() @@ -104,7 +106,11 @@ class TestSagemakerTransform: def test_map_mistral_params(self): """Test that parameters are correctly mapped""" - test_params = {"temperature": 0.7, "max_tokens": 200, "max_completion_tokens": 256} + test_params = { + "temperature": 0.7, + "max_tokens": 200, + "max_completion_tokens": 256, + } result = self.config.map_openai_params( non_default_params=test_params, @@ -118,7 +124,10 @@ class TestSagemakerTransform: def test_mistral_max_tokens_backward_compat(self): """Test that parameters are correctly mapped""" - test_params = {"temperature": 0.7, "max_tokens": 200,} + test_params = { + "temperature": 0.7, + "max_tokens": 200, + } result = self.config.map_openai_params( non_default_params=test_params, diff --git a/tests/litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py similarity index 100% rename from tests/litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py rename to tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py diff --git a/tests/litellm/llms/vertex_ai/multimodal_embeddings/test_vertex_ai_multimodal_embedding_transformation.py b/tests/test_litellm/llms/vertex_ai/multimodal_embeddings/test_vertex_ai_multimodal_embedding_transformation.py similarity index 100% rename from tests/litellm/llms/vertex_ai/multimodal_embeddings/test_vertex_ai_multimodal_embedding_transformation.py rename to tests/test_litellm/llms/vertex_ai/multimodal_embeddings/test_vertex_ai_multimodal_embedding_transformation.py diff --git a/tests/litellm/llms/vertex_ai/test_http_status_201.py b/tests/test_litellm/llms/vertex_ai/test_http_status_201.py similarity index 90% rename from tests/litellm/llms/vertex_ai/test_http_status_201.py rename to tests/test_litellm/llms/vertex_ai/test_http_status_201.py index 6d13eef6049..3d6c5b538e9 100644 --- a/tests/litellm/llms/vertex_ai/test_http_status_201.py +++ b/tests/test_litellm/llms/vertex_ai/test_http_status_201.py @@ -1,34 +1,36 @@ import json import unittest -from unittest.mock import patch, MagicMock +from unittest.mock import MagicMock, patch +from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexAIError, make_call, make_sync_call, - VertexAIError, ) -from litellm.llms.custom_httpx.http_handler import HTTPHandler class TestVertexAIHTTPStatus201(unittest.TestCase): def setUp(self): # Setup mock messages self.messages = [{"role": "user", "content": "Hello, how are you?"}] - + # Setup mock data self.mock_data = json.dumps({"messages": self.messages}) - + # Setup mock headers self.mock_headers = {"Content-Type": "application/json"} - + # Setup mock model self.mock_model = "gemini-pro" - + # Setup mock logging object self.mock_logging_obj = MagicMock() self.mock_logging_obj.post_call = MagicMock() - @patch("litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.get_async_httpx_client") + @patch( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.get_async_httpx_client" + ) async def test_async_http_status_201(self, mock_get_client): """Test that async make_call handles HTTP 201 status code correctly""" # Create a mock response with status code 201 @@ -36,13 +38,13 @@ class TestVertexAIHTTPStatus201(unittest.TestCase): mock_response.status_code = 201 mock_response.aiter_lines = MagicMock() mock_response.aiter_lines.return_value = ["test response"] - + # Setup mock client mock_client = MagicMock() mock_client.post = MagicMock() mock_client.post.return_value = mock_response mock_get_client.return_value = mock_client - + # Call the make_call function result = await make_call( client=None, @@ -51,15 +53,15 @@ class TestVertexAIHTTPStatus201(unittest.TestCase): data=self.mock_data, model=self.mock_model, messages=self.messages, - logging_obj=self.mock_logging_obj + logging_obj=self.mock_logging_obj, ) - + # Assert that the post method was called mock_client.post.assert_called_once() - + # Assert that no error was raised for status code 201 self.assertIsNotNone(result) - + # Verify logging was called self.mock_logging_obj.post_call.assert_called_once() @@ -72,7 +74,7 @@ class TestVertexAIHTTPStatus201(unittest.TestCase): mock_response.iter_lines = MagicMock() mock_response.iter_lines.return_value = ["test response"] mock_post.return_value = mock_response - + # Call the make_sync_call function result = make_sync_call( client=None, @@ -82,12 +84,12 @@ class TestVertexAIHTTPStatus201(unittest.TestCase): data=self.mock_data, model=self.mock_model, messages=self.messages, - logging_obj=self.mock_logging_obj + logging_obj=self.mock_logging_obj, ) - + # Assert that no error was raised for status code 201 self.assertIsNotNone(result) - + # Verify logging was called self.mock_logging_obj.post_call.assert_called_once() @@ -100,7 +102,7 @@ class TestVertexAIHTTPStatus201(unittest.TestCase): mock_response.read = MagicMock(return_value=b"Bad Request") mock_response.headers = {} mock_post.return_value = mock_response - + # Call the make_sync_call function and expect an error with self.assertRaises(VertexAIError) as context: make_sync_call( @@ -111,12 +113,12 @@ class TestVertexAIHTTPStatus201(unittest.TestCase): data=self.mock_data, model=self.mock_model, messages=self.messages, - logging_obj=self.mock_logging_obj + logging_obj=self.mock_logging_obj, ) - + # Assert that the error has the correct status code self.assertEqual(context.exception.status_code, 400) if __name__ == "__main__": - unittest.main() \ No newline at end of file + unittest.main() diff --git a/tests/litellm/llms/vertex_ai/test_vertex.py b/tests/test_litellm/llms/vertex_ai/test_vertex.py similarity index 99% rename from tests/litellm/llms/vertex_ai/test_vertex.py rename to tests/test_litellm/llms/vertex_ai/test_vertex.py index 05930769f4b..1b9de4d0fd9 100644 --- a/tests/litellm/llms/vertex_ai/test_vertex.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex.py @@ -1,9 +1,7 @@ import base64 -import numpy as np import json import os import sys -import traceback from dotenv import load_dotenv @@ -18,6 +16,7 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path import pytest + import litellm from litellm import get_optional_params from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_image @@ -31,6 +30,7 @@ def encode_image_to_base64(image_path): def test_completion_pydantic_obj_2(): from pydantic import BaseModel + from litellm.llms.custom_httpx.http_handler import HTTPHandler litellm.set_verbose = True @@ -110,11 +110,10 @@ def test_completion_pydantic_obj_2(): def test_build_vertex_schema(): - from litellm.llms.vertex_ai.common_utils import ( - _build_vertex_schema, - ) import json + from litellm.llms.vertex_ai.common_utils import _build_vertex_schema + schema = { "type": "object", "my-random-key": "my-random-value", @@ -149,7 +148,6 @@ def test_build_vertex_schema(): ], ) def test_vertex_tool_params(tools, key): - optional_params = get_optional_params( model="gemini-1.5-pro", custom_llm_provider="vertex_ai", @@ -1124,7 +1122,6 @@ def test_logprobs(): mock_response.json.return_value = response_body with patch.object(client, "post", return_value=mock_response): - resp = litellm.completion( model="gemini/gemini-1.5-flash-002", messages=[ @@ -1140,9 +1137,7 @@ def test_logprobs(): def test_process_gemini_image(): """Test the _process_gemini_image function for different image sources""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _process_gemini_image, - ) + from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_image from litellm.types.llms.vertex_ai import FileDataType # Test GCS URI @@ -1266,9 +1261,10 @@ def test_vertex_embedding_url(model, expected_url): assert endpoint == "predict" -import pytest from unittest.mock import Mock, patch +import pytest + # Add these fixtures below existing fixtures @pytest.fixture @@ -1435,7 +1431,10 @@ def test_vertex_parallel_tool_calls_false_multiple_tools_error(): tools=tools, parallel_tool_calls=False, ) - assert "`parallel_tool_calls=False` is not supported when multiple tools are provided" in str(excinfo.value) + assert ( + "`parallel_tool_calls=False` is not supported when multiple tools are provided" + in str(excinfo.value) + ) # works when specified as "functions" with pytest.raises(litellm.utils.UnsupportedParamsError) as excinfo: @@ -1445,7 +1444,10 @@ def test_vertex_parallel_tool_calls_false_multiple_tools_error(): functions=tools, parallel_tool_calls=False, ) - assert "`parallel_tool_calls=False` is not supported when multiple tools are provided" in str(excinfo.value) + assert ( + "`parallel_tool_calls=False` is not supported when multiple tools are provided" + in str(excinfo.value) + ) def test_vertex_parallel_tool_calls_false_single_tool(): diff --git a/tests/litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py similarity index 99% rename from tests/litellm/llms/vertex_ai/test_vertex_ai_common_utils.py rename to tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py index fca89650a53..8ae49a2aa73 100644 --- a/tests/litellm/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -13,11 +13,11 @@ sys.path.insert( import litellm from litellm.llms.vertex_ai.common_utils import ( + _get_vertex_url, convert_anyof_null_to_nullable, get_vertex_location_from_url, get_vertex_project_id_from_url, set_schema_property_ordering, - _get_vertex_url ) @@ -518,6 +518,7 @@ def test_vertex_ai_complex_response_schema(): assert "additionalProperties" not in type3 assert "additionalProperties" not in type3_prop3_items + @pytest.mark.parametrize( "stream, expected_endpoint_suffix", [ @@ -537,7 +538,10 @@ def test_get_vertex_url_global_region(stream, expected_endpoint_suffix): # Mock litellm.VertexGeminiConfig.get_model_for_vertex_ai_url to return model as is # as we are not testing that part here, just the URL construction - with patch("litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model): + with patch( + "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", + side_effect=lambda model: model, + ): url, endpoint = _get_vertex_url( mode=mode, model=model, @@ -548,7 +552,7 @@ def test_get_vertex_url_global_region(stream, expected_endpoint_suffix): ) expected_url_base = f"https://aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/global/publishers/google/models/{model}" - + if stream: expected_endpoint = "streamGenerateContent" expected_url = f"{expected_url_base}:{expected_endpoint}?alt=sse" @@ -556,6 +560,5 @@ def test_get_vertex_url_global_region(stream, expected_endpoint_suffix): expected_endpoint = "generateContent" expected_url = f"{expected_url_base}:{expected_endpoint}" - assert endpoint == expected_endpoint assert url == expected_url diff --git a/tests/litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py similarity index 100% rename from tests/litellm/llms/vertex_ai/test_vertex_llm_base.py rename to tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py diff --git a/tests/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py similarity index 100% rename from tests/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py rename to tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py diff --git a/tests/litellm/log.txt b/tests/test_litellm/log.txt similarity index 100% rename from tests/litellm/log.txt rename to tests/test_litellm/log.txt diff --git a/tests/test_litellm/proxy/anthropic_endpoints/__init__.py b/tests/test_litellm/proxy/anthropic_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/litellm/proxy/anthropic_endpoints/test_endpoints.py b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py similarity index 65% rename from tests/litellm/proxy/anthropic_endpoints/test_endpoints.py rename to tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py index 2dbdf345042..6c2d25c6752 100644 --- a/tests/litellm/proxy/anthropic_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py @@ -24,38 +24,49 @@ class TestAnthropicEndpoints(unittest.TestCase): {"type": "content_block_delta", "delta": {"text": "more data"}}, "text chunk data again", ] - + mock_user_api_key_dict = MagicMock() mock_request_data = {} mock_proxy_logging_obj = MagicMock() - mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["response"]) - + mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock( + side_effect=lambda **kwargs: kwargs["response"] + ) + # Configure safe_dumps to return a properly formatted JSON string mock_safe_dumps.side_effect = lambda chunk: json.dumps(chunk) - + # Execute - result = [chunk async for chunk in async_data_generator_anthropic( - response=mock_response, - user_api_key_dict=mock_user_api_key_dict, - request_data=mock_request_data, - proxy_logging_obj=mock_proxy_logging_obj, - )] - + result = [ + chunk + async for chunk in async_data_generator_anthropic( + response=mock_response, + user_api_key_dict=mock_user_api_key_dict, + request_data=mock_request_data, + proxy_logging_obj=mock_proxy_logging_obj, + ) + ] + # Verify expected_result = [ 'data: {"type": "message_start", "message": {"id": "msg_123"}}\n\n', - 'text chunk data', + "text chunk data", 'data: {"type": "content_block_delta", "delta": {"text": "more data"}}\n\n', - 'text chunk data again', + "text chunk data again", ] - + self.assertEqual(result, expected_result) - + # Assert safe_dumps was called for dictionary objects - mock_safe_dumps.assert_any_call({"type": "message_start", "message": {"id": "msg_123"}}) - mock_safe_dumps.assert_any_call({"type": "content_block_delta", "delta": {"text": "more data"}}) - assert mock_safe_dumps.call_count == 2 # Called twice, once for each dict object + mock_safe_dumps.assert_any_call( + {"type": "message_start", "message": {"id": "msg_123"}} + ) + mock_safe_dumps.assert_any_call( + {"type": "content_block_delta", "delta": {"text": "more data"}} + ) + assert ( + mock_safe_dumps.call_count == 2 + ) # Called twice, once for each dict object if __name__ == "__main__": - unittest.main() \ No newline at end of file + unittest.main() diff --git a/tests/litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py similarity index 100% rename from tests/litellm/proxy/auth/test_auth_checks.py rename to tests/test_litellm/proxy/auth/test_auth_checks.py diff --git a/tests/litellm/proxy/auth/test_auth_exception_handler.py b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py similarity index 100% rename from tests/litellm/proxy/auth/test_auth_exception_handler.py rename to tests/test_litellm/proxy/auth/test_auth_exception_handler.py diff --git a/tests/litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py similarity index 100% rename from tests/litellm/proxy/auth/test_handle_jwt.py rename to tests/test_litellm/proxy/auth/test_handle_jwt.py diff --git a/tests/litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py similarity index 100% rename from tests/litellm/proxy/auth/test_user_api_key_auth.py rename to tests/test_litellm/proxy/auth/test_user_api_key_auth.py diff --git a/tests/test_litellm/proxy/client/cli/__init__.py b/tests/test_litellm/proxy/client/cli/__init__.py new file mode 100644 index 00000000000..14e8836e483 --- /dev/null +++ b/tests/test_litellm/proxy/client/cli/__init__.py @@ -0,0 +1 @@ +"""Tests for the LiteLLM Proxy Client CLI package.""" diff --git a/tests/litellm/proxy/client/cli/test_chat_commands.py b/tests/test_litellm/proxy/client/cli/test_chat_commands.py similarity index 99% rename from tests/litellm/proxy/client/cli/test_chat_commands.py rename to tests/test_litellm/proxy/client/cli/test_chat_commands.py index 796229ccc45..0b785f4962c 100644 --- a/tests/litellm/proxy/client/cli/test_chat_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_chat_commands.py @@ -1,5 +1,5 @@ import json -from unittest.mock import patch, MagicMock +from unittest.mock import MagicMock, patch import pytest import requests @@ -238,4 +238,4 @@ def test_chat_completions_all_parameters(cli_runner, mock_chat_client): presence_penalty=0.5, frequency_penalty=0.5, user="test-user", - ) \ No newline at end of file + ) diff --git a/tests/litellm/proxy/client/cli/test_credentials_commands.py b/tests/test_litellm/proxy/client/cli/test_credentials_commands.py similarity index 94% rename from tests/litellm/proxy/client/cli/test_credentials_commands.py rename to tests/test_litellm/proxy/client/cli/test_credentials_commands.py index de2bdec5ce4..cae9e79d797 100644 --- a/tests/litellm/proxy/client/cli/test_credentials_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_credentials_commands.py @@ -1,5 +1,5 @@ import json -from unittest.mock import patch, MagicMock +from unittest.mock import MagicMock, patch import pytest import requests @@ -10,7 +10,9 @@ from litellm.proxy.client.cli.main import cli @pytest.fixture def mock_credentials_client(): - with patch("litellm.proxy.client.cli.commands.credentials.CredentialsManagementClient") as mock: + with patch( + "litellm.proxy.client.cli.commands.credentials.CredentialsManagementClient" + ) as mock: yield mock @@ -128,7 +130,9 @@ def test_create_credential_http_error(cli_runner, mock_credentials_client): mock_error_response = MagicMock() mock_error_response.status_code = 400 mock_error_response.json.return_value = {"error": "Invalid request"} - mock_instance.create.side_effect = requests.exceptions.HTTPError(response=mock_error_response) + mock_instance.create.side_effect = requests.exceptions.HTTPError( + response=mock_error_response + ) # Run command result = cli_runner.invoke( @@ -172,7 +176,9 @@ def test_delete_credential_http_error(cli_runner, mock_credentials_client): mock_error_response = MagicMock() mock_error_response.status_code = 404 mock_error_response.json.return_value = {"error": "Credential not found"} - mock_instance.delete.side_effect = requests.exceptions.HTTPError(response=mock_error_response) + mock_instance.delete.side_effect = requests.exceptions.HTTPError( + response=mock_error_response + ) # Run command result = cli_runner.invoke(cli, ["credentials", "delete", "test-cred"]) @@ -199,4 +205,4 @@ def test_get_credential_success(cli_runner, mock_credentials_client): assert result.exit_code == 0 output_data = json.loads(result.output) assert output_data == mock_response - mock_instance.get.assert_called_once_with("test-cred") \ No newline at end of file + mock_instance.get.assert_called_once_with("test-cred") diff --git a/tests/litellm/proxy/client/cli/test_global_options.py b/tests/test_litellm/proxy/client/cli/test_global_options.py similarity index 74% rename from tests/litellm/proxy/client/cli/test_global_options.py rename to tests/test_litellm/proxy/client/cli/test_global_options.py index 9a2f64a2780..b14bfdf92af 100644 --- a/tests/litellm/proxy/client/cli/test_global_options.py +++ b/tests/test_litellm/proxy/client/cli/test_global_options.py @@ -1,10 +1,12 @@ # stdlib imports -from litellm.proxy.client.cli import cli -from litellm._version import version as litellm_version -from click.testing import CliRunner -import pytest -from unittest.mock import patch import os +from unittest.mock import patch + +import pytest +from click.testing import CliRunner + +from litellm._version import version as litellm_version +from litellm.proxy.client.cli import cli @pytest.fixture @@ -14,8 +16,10 @@ def cli_runner(): def test_cli_version_flag(cli_runner): """Test that --version prints the correct version, server URL, and server version, and exits successfully""" - with patch("litellm.proxy.client.health.HealthManagementClient.get_server_version", return_value="1.2.3"), \ - patch.dict(os.environ, {"LITELLM_PROXY_URL": "http://localhost:4000"}): + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ), patch.dict(os.environ, {"LITELLM_PROXY_URL": "http://localhost:4000"}): result = cli_runner.invoke(cli, ["--version"]) assert result.exit_code == 0 assert f"LiteLLM Proxy CLI Version: {litellm_version}" in result.output @@ -25,8 +29,10 @@ def test_cli_version_flag(cli_runner): def test_cli_version_command(cli_runner): """Test that 'version' command prints the correct version, server URL, and server version, and exits successfully""" - with patch("litellm.proxy.client.health.HealthManagementClient.get_server_version", return_value="1.2.3"), \ - patch.dict(os.environ, {"LITELLM_PROXY_URL": "http://localhost:4000"}): + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ), patch.dict(os.environ, {"LITELLM_PROXY_URL": "http://localhost:4000"}): result = cli_runner.invoke(cli, ["version"]) assert result.exit_code == 0 assert f"LiteLLM Proxy CLI Version: {litellm_version}" in result.output diff --git a/tests/litellm/proxy/client/cli/test_keys_commands.py b/tests/test_litellm/proxy/client/cli/test_keys_commands.py similarity index 71% rename from tests/litellm/proxy/client/cli/test_keys_commands.py rename to tests/test_litellm/proxy/client/cli/test_keys_commands.py index d3f1dcfb9a3..aaaedc458f4 100644 --- a/tests/litellm/proxy/client/cli/test_keys_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_keys_commands.py @@ -15,19 +15,35 @@ def cli_runner(): @pytest.fixture(autouse=True) def mock_env(): - with patch.dict(os.environ, {"LITELLM_PROXY_URL": "http://localhost:4000", "LITELLM_PROXY_API_KEY": "sk-test"}): + with patch.dict( + os.environ, + { + "LITELLM_PROXY_URL": "http://localhost:4000", + "LITELLM_PROXY_API_KEY": "sk-test", + }, + ): yield @pytest.fixture def mock_keys_client(): - with patch("litellm.proxy.client.cli.commands.keys.KeysManagementClient") as MockClient: + with patch( + "litellm.proxy.client.cli.commands.keys.KeysManagementClient" + ) as MockClient: yield MockClient def test_keys_list_json_format(mock_keys_client, cli_runner): mock_keys_client.return_value.list.return_value = { - "keys": [{"token": "abc123", "key_alias": "alias1", "user_id": "u1", "team_id": "t1", "spend": 10.0}] + "keys": [ + { + "token": "abc123", + "key_alias": "alias1", + "user_id": "u1", + "team_id": "t1", + "spend": 10.0, + } + ] } result = cli_runner.invoke(cli, ["keys", "list", "--format", "json"]) assert result.exit_code == 0 @@ -39,7 +55,15 @@ def test_keys_list_json_format(mock_keys_client, cli_runner): def test_keys_list_table_format(mock_keys_client, cli_runner): mock_keys_client.return_value.list.return_value = { - "keys": [{"token": "abc123", "key_alias": "alias1", "user_id": "u1", "team_id": "t1", "spend": 10.0}] + "keys": [ + { + "token": "abc123", + "key_alias": "alias1", + "user_id": "u1", + "team_id": "t1", + "spend": 10.0, + } + ] } result = cli_runner.invoke(cli, ["keys", "list"]) assert result.exit_code == 0 @@ -53,15 +77,23 @@ def test_keys_list_table_format(mock_keys_client, cli_runner): def test_keys_generate_success(mock_keys_client, cli_runner): - mock_keys_client.return_value.generate.return_value = {"key": "new-key", "spend": 100.0} - result = cli_runner.invoke(cli, ["keys", "generate", "--models", "gpt-4", "--spend", "100"]) + mock_keys_client.return_value.generate.return_value = { + "key": "new-key", + "spend": 100.0, + } + result = cli_runner.invoke( + cli, ["keys", "generate", "--models", "gpt-4", "--spend", "100"] + ) assert result.exit_code == 0 assert "new-key" in result.output mock_keys_client.return_value.generate.assert_called_once() def test_keys_delete_success(mock_keys_client, cli_runner): - mock_keys_client.return_value.delete.return_value = {"status": "success", "deleted_keys": ["abc123"]} + mock_keys_client.return_value.delete.return_value = { + "status": "success", + "deleted_keys": ["abc123"], + } result = cli_runner.invoke(cli, ["keys", "delete", "--keys", "abc123"]) assert result.exit_code == 0 assert "success" in result.output diff --git a/tests/litellm/proxy/client/cli/test_models_commands.py b/tests/test_litellm/proxy/client/cli/test_models_commands.py similarity index 84% rename from tests/litellm/proxy/client/cli/test_models_commands.py rename to tests/test_litellm/proxy/client/cli/test_models_commands.py index f05a20ee53d..16744f4a895 100644 --- a/tests/litellm/proxy/client/cli/test_models_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_models_commands.py @@ -4,16 +4,17 @@ import os import time from unittest.mock import patch +import pytest + # third party imports from click.testing import CliRunner -import pytest # local imports from litellm.proxy.client.cli import cli from litellm.proxy.client.cli.commands.models import ( - format_timestamp, - format_iso_datetime_str, format_cost_per_1k_tokens, + format_iso_datetime_str, + format_timestamp, ) @@ -47,8 +48,18 @@ def mock_env(): def mock_models_list(mock_client): """Fixture to set up common mocking pattern for models list tests""" mock_client.return_value.models.list.return_value = [ - {"id": "model-123", "object": "model", "created": 1699848889, "owned_by": "organization-123"}, - {"id": "model-456", "object": "model", "created": 1699848890, "owned_by": "organization-456"}, + { + "id": "model-123", + "object": "model", + "created": 1699848889, + "owned_by": "organization-123", + }, + { + "id": "model-456", + "object": "model", + "created": 1699848890, + "owned_by": "organization-456", + }, ] mock_client.assert_not_called() # Ensure clean slate @@ -107,7 +118,9 @@ def test_models_list_json_format(mock_models_list, cli_runner): assert output_data == mock_models_list.return_value.models.list.return_value # Verify the client was called correctly - mock_models_list.assert_called_once_with(base_url="http://localhost:4000", api_key="sk-test") + mock_models_list.assert_called_once_with( + base_url="http://localhost:4000", api_key="sk-test" + ) mock_models_list.return_value.models.list.assert_called_once() @@ -129,7 +142,9 @@ def test_models_list_table_format(mock_models_list, cli_runner): assert format_timestamp(1699848889) in result.output # Verify the client was called correctly - mock_models_list.assert_called_once_with(base_url="http://localhost:4000", api_key="sk-test") + mock_models_list.assert_called_once_with( + base_url="http://localhost:4000", api_key="sk-test" + ) mock_models_list.return_value.models.list.assert_called_once() @@ -180,7 +195,9 @@ def test_models_list_error_handling(mock_client, cli_runner): assert "API Error" in str(result.exception) # Verify the client was created with env var values - mock_client.assert_called_once_with(base_url="http://localhost:4000", api_key="sk-test") + mock_client.assert_called_once_with( + base_url="http://localhost:4000", api_key="sk-test" + ) def test_models_info_json_format(mock_models_info, cli_runner): @@ -196,7 +213,9 @@ def test_models_info_json_format(mock_models_info, cli_runner): assert output_data == mock_models_info.return_value.models.info.return_value # Verify the client was called correctly with env var values - mock_models_info.assert_called_once_with(base_url="http://localhost:4000", api_key="sk-test") + mock_models_info.assert_called_once_with( + base_url="http://localhost:4000", api_key="sk-test" + ) mock_models_info.return_value.models.info.assert_called_once() @@ -220,7 +239,9 @@ def test_models_info_table_format(mock_models_info, cli_runner): assert "843000" not in result.output # Verify the client was called correctly with env var values - mock_models_info.assert_called_once_with(base_url="http://localhost:4000", api_key="sk-test") + mock_models_info.assert_called_once_with( + base_url="http://localhost:4000", api_key="sk-test" + ) mock_models_info.return_value.models.info.assert_called_once() @@ -229,10 +250,26 @@ def test_models_import_only_models_matching_regex(tmp_path, mock_client, cli_run # Prepare a YAML file with a mix of models yaml_content = { "model_list": [ - {"model_name": "gpt-4-model", "litellm_params": {"model": "gpt-4"}, "model_info": {"id": "id-1"}}, - {"model_name": "gpt-3.5-model", "litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "id-2"}}, - {"model_name": "llama2-model", "litellm_params": {"model": "llama2"}, "model_info": {"id": "id-3"}}, - {"model_name": "other-model", "litellm_params": {"model": "other"}, "model_info": {"id": "id-4"}}, + { + "model_name": "gpt-4-model", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "id-1"}, + }, + { + "model_name": "gpt-3.5-model", + "litellm_params": {"model": "gpt-3.5-turbo"}, + "model_info": {"id": "id-2"}, + }, + { + "model_name": "llama2-model", + "litellm_params": {"model": "llama2"}, + "model_info": {"id": "id-3"}, + }, + { + "model_name": "other-model", + "litellm_params": {"model": "other"}, + "model_info": {"id": "id-4"}, + }, ] } import yaml as pyyaml @@ -245,7 +282,9 @@ def test_models_import_only_models_matching_regex(tmp_path, mock_client, cli_run mock_new = mock_client.return_value.models.new # Only match models containing 'gpt' in their litellm_params.model - result = cli_runner.invoke(cli, ["models", "import", str(yaml_file), "--only-models-matching-regex", "gpt"]) + result = cli_runner.invoke( + cli, ["models", "import", str(yaml_file), "--only-models-matching-regex", "gpt"] + ) # Should succeed assert result.exit_code == 0 @@ -259,7 +298,9 @@ def test_models_import_only_models_matching_regex(tmp_path, mock_client, cli_run assert "gpt-4".split("-")[0] in result.output or "gpt" in result.output -def test_models_import_only_access_groups_matching_regex(tmp_path, mock_client, cli_runner): +def test_models_import_only_access_groups_matching_regex( + tmp_path, mock_client, cli_runner +): """Test the --only-access-groups-matching-regex option for models import command""" # Prepare a YAML file with a mix of models yaml_content = { @@ -267,7 +308,10 @@ def test_models_import_only_access_groups_matching_regex(tmp_path, mock_client, { "model_name": "gpt-4-model", "litellm_params": {"model": "gpt-4"}, - "model_info": {"id": "id-1", "access_groups": ["beta-models", "prod-models"]}, + "model_info": { + "id": "id-1", + "access_groups": ["beta-models", "prod-models"], + }, }, { "model_name": "gpt-3.5-model", @@ -301,7 +345,16 @@ def test_models_import_only_access_groups_matching_regex(tmp_path, mock_client, mock_new = mock_client.return_value.models.new # Only match models with access_groups containing 'beta' - result = cli_runner.invoke(cli, ["models", "import", str(yaml_file), "--only-access-groups-matching-regex", "beta"]) + result = cli_runner.invoke( + cli, + [ + "models", + "import", + str(yaml_file), + "--only-access-groups-matching-regex", + "beta", + ], + ) # Should succeed assert result.exit_code == 0 diff --git a/tests/litellm/proxy/client/cli/test_users_commands.py b/tests/test_litellm/proxy/client/cli/test_users_commands.py similarity index 62% rename from tests/litellm/proxy/client/cli/test_users_commands.py rename to tests/test_litellm/proxy/client/cli/test_users_commands.py index 3489855e81f..9b291c4ec9d 100644 --- a/tests/litellm/proxy/client/cli/test_users_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_users_commands.py @@ -1,26 +1,50 @@ +from unittest.mock import patch + import pytest from click.testing import CliRunner -from unittest.mock import patch + from litellm.proxy.client.cli import cli + @pytest.fixture def cli_runner(): return CliRunner() + @pytest.fixture(autouse=True) def mock_env(): - with patch.dict("os.environ", {"LITELLM_PROXY_URL": "http://localhost:4000", "LITELLM_PROXY_API_KEY": "sk-test"}): + with patch.dict( + "os.environ", + { + "LITELLM_PROXY_URL": "http://localhost:4000", + "LITELLM_PROXY_API_KEY": "sk-test", + }, + ): yield + @pytest.fixture def mock_users_client(): - with patch("litellm.proxy.client.cli.commands.users.UsersManagementClient") as MockClient: + with patch( + "litellm.proxy.client.cli.commands.users.UsersManagementClient" + ) as MockClient: yield MockClient + def test_users_list(cli_runner, mock_users_client): mock_users_client.return_value.list_users.return_value = [ - {"user_id": "u1", "user_email": "a@b.com", "user_role": "internal_user", "teams": ["t1"]}, - {"user_id": "u2", "user_email": "b@b.com", "user_role": "proxy_admin", "teams": ["t2", "t3"]}, + { + "user_id": "u1", + "user_email": "a@b.com", + "user_role": "internal_user", + "teams": ["t1"], + }, + { + "user_id": "u2", + "user_email": "b@b.com", + "user_role": "proxy_admin", + "teams": ["t2", "t3"], + }, ] result = cli_runner.invoke(cli, ["users", "list"]) assert result.exit_code == 0 @@ -30,25 +54,36 @@ def test_users_list(cli_runner, mock_users_client): assert "t3" in result.output mock_users_client.return_value.list_users.assert_called_once() + def test_users_get(cli_runner, mock_users_client): - mock_users_client.return_value.get_user.return_value = {"user_id": "u1", "user_email": "a@b.com"} + mock_users_client.return_value.get_user.return_value = { + "user_id": "u1", + "user_email": "a@b.com", + } result = cli_runner.invoke(cli, ["users", "get", "--id", "u1"]) assert result.exit_code == 0 assert '"user_id": "u1"' in result.output assert '"user_email": "a@b.com"' in result.output mock_users_client.return_value.get_user.assert_called_once_with(user_id="u1") + def test_users_create(cli_runner, mock_users_client): - mock_users_client.return_value.create_user.return_value = {"user_id": "u1", "user_email": "a@b.com"} - result = cli_runner.invoke(cli, ["users", "create", "--email", "a@b.com", "--role", "internal_user"]) + mock_users_client.return_value.create_user.return_value = { + "user_id": "u1", + "user_email": "a@b.com", + } + result = cli_runner.invoke( + cli, ["users", "create", "--email", "a@b.com", "--role", "internal_user"] + ) assert result.exit_code == 0 assert '"user_id": "u1"' in result.output assert '"user_email": "a@b.com"' in result.output mock_users_client.return_value.create_user.assert_called_once() + def test_users_delete(cli_runner, mock_users_client): mock_users_client.return_value.delete_user.return_value = {"deleted": 1} result = cli_runner.invoke(cli, ["users", "delete", "u1", "u2"]) assert result.exit_code == 0 assert '"deleted": 1' in result.output - mock_users_client.return_value.delete_user.assert_called_once_with(["u1", "u2"]) \ No newline at end of file + mock_users_client.return_value.delete_user.assert_called_once_with(["u1", "u2"]) diff --git a/tests/litellm/proxy/client/test_chat.py b/tests/test_litellm/proxy/client/test_chat.py similarity index 81% rename from tests/litellm/proxy/client/test_chat.py rename to tests/test_litellm/proxy/client/test_chat.py index 95646e6a26b..2fd5a4a26ff 100644 --- a/tests/litellm/proxy/client/test_chat.py +++ b/tests/test_litellm/proxy/client/test_chat.py @@ -1,5 +1,6 @@ import pytest import requests + from litellm.proxy.client.chat import ChatClient from litellm.proxy.client.exceptions import UnauthorizedError @@ -53,7 +54,11 @@ def test_client_without_api_key(base_url): def test_completions_request_creation(client, base_url, api_key, sample_messages): """Test that completions creates a request with correct URL, headers, and body""" request = client.completions( - model="gpt-4", messages=sample_messages, temperature=0.7, max_tokens=100, return_request=True + model="gpt-4", + messages=sample_messages, + temperature=0.7, + max_tokens=100, + return_request=True, ) # Check request method and URL @@ -65,12 +70,19 @@ def test_completions_request_creation(client, base_url, api_key, sample_messages assert request.headers["Authorization"] == f"Bearer {api_key}" # Check request body - assert request.json == {"model": "gpt-4", "messages": sample_messages, "temperature": 0.7, "max_tokens": 100} + assert request.json == { + "model": "gpt-4", + "messages": sample_messages, + "temperature": 0.7, + "max_tokens": 100, + } def test_completions_minimal_request(client, sample_messages): """Test that completions works with only required parameters""" - request = client.completions(model="gpt-4", messages=sample_messages, return_request=True) + request = client.completions( + model="gpt-4", messages=sample_messages, return_request=True + ) # Check request body has only required fields assert request.json == {"model": "gpt-4", "messages": sample_messages} @@ -115,7 +127,10 @@ def test_completions_mock_response(client, sample_messages, requests_mock): "usage": {"prompt_tokens": 13, "completion_tokens": 7, "total_tokens": 20}, "choices": [ { - "message": {"role": "assistant", "content": "Hello! How can I help you today?"}, + "message": { + "role": "assistant", + "content": "Hello! How can I help you today?", + }, "finish_reason": "stop", "index": 0, } @@ -128,13 +143,20 @@ def test_completions_mock_response(client, sample_messages, requests_mock): response = client.completions(model="gpt-4", messages=sample_messages) assert response == mock_response - assert response["choices"][0]["message"]["content"] == "Hello! How can I help you today?" + assert ( + response["choices"][0]["message"]["content"] + == "Hello! How can I help you today?" + ) def test_completions_unauthorized_error(client, sample_messages, requests_mock): """Test that completions raises UnauthorizedError for 401 responses""" # Mock a 401 response - requests_mock.post(f"{client._base_url}/chat/completions", status_code=401, json={"error": "Unauthorized"}) + requests_mock.post( + f"{client._base_url}/chat/completions", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.completions(model="gpt-4", messages=sample_messages) @@ -143,7 +165,11 @@ def test_completions_unauthorized_error(client, sample_messages, requests_mock): def test_completions_other_errors(client, sample_messages, requests_mock): """Test that completions raises HTTPError for other error responses""" # Mock a 500 response - requests_mock.post(f"{client._base_url}/chat/completions", status_code=500, json={"error": "Internal Server Error"}) + requests_mock.post( + f"{client._base_url}/chat/completions", + status_code=500, + json={"error": "Internal Server Error"}, + ) with pytest.raises(requests.exceptions.HTTPError) as exc_info: client.completions(model="gpt-4", messages=sample_messages) diff --git a/tests/litellm/proxy/client/test_client.py b/tests/test_litellm/proxy/client/test_client.py similarity index 97% rename from tests/litellm/proxy/client/test_client.py rename to tests/test_litellm/proxy/client/test_client.py index 9806237992e..267e9f074a4 100644 --- a/tests/litellm/proxy/client/test_client.py +++ b/tests/test_litellm/proxy/client/test_client.py @@ -1,7 +1,8 @@ import pytest -from litellm.proxy.client import Client, ModelsManagementClient, ChatClient -from litellm.proxy.client.keys import KeysManagementClient + +from litellm.proxy.client import ChatClient, Client, ModelsManagementClient from litellm.proxy.client.http_client import HTTPClient +from litellm.proxy.client.keys import KeysManagementClient @pytest.fixture diff --git a/tests/litellm/proxy/client/test_credentials.py b/tests/test_litellm/proxy/client/test_credentials.py similarity index 78% rename from tests/litellm/proxy/client/test_credentials.py rename to tests/test_litellm/proxy/client/test_credentials.py index f5af522c598..ef0cc9998c8 100644 --- a/tests/litellm/proxy/client/test_credentials.py +++ b/tests/test_litellm/proxy/client/test_credentials.py @@ -1,5 +1,6 @@ import pytest import requests + from litellm.proxy.client.credentials import CredentialsManagementClient from litellm.proxy.client.exceptions import UnauthorizedError @@ -77,7 +78,11 @@ def test_list_mock_response(client, requests_mock): def test_list_unauthorized_error(client, requests_mock): """Test that list raises UnauthorizedError for 401 responses""" - requests_mock.get(f"{client._base_url}/credentials", status_code=401, json={"error": "Unauthorized"}) + requests_mock.get( + f"{client._base_url}/credentials", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.list() @@ -88,7 +93,10 @@ def test_create_request(client, base_url, api_key): request = client.create( credential_name="azure1", credential_info={"api_type": "azure"}, - credential_values={"api_key": "sk-123", "api_base": "https://example.azure.openai.com"}, + credential_values={ + "api_key": "sk-123", + "api_base": "https://example.azure.openai.com", + }, return_request=True, ) @@ -99,31 +107,47 @@ def test_create_request(client, base_url, api_key): assert request.json == { "credential_name": "azure1", "credential_info": {"api_type": "azure"}, - "credential_values": {"api_key": "sk-123", "api_base": "https://example.azure.openai.com"}, + "credential_values": { + "api_key": "sk-123", + "api_base": "https://example.azure.openai.com", + }, } def test_create_mock_response(client, requests_mock): """Test create with a mocked successful response""" - mock_response = {"credential_name": "azure1", "credential_info": {"api_type": "azure"}, "status": "success"} + mock_response = { + "credential_name": "azure1", + "credential_info": {"api_type": "azure"}, + "status": "success", + } requests_mock.post(f"{client._base_url}/credentials", json=mock_response) response = client.create( credential_name="azure1", credential_info={"api_type": "azure"}, - credential_values={"api_key": "sk-123", "api_base": "https://example.azure.openai.com"}, + credential_values={ + "api_key": "sk-123", + "api_base": "https://example.azure.openai.com", + }, ) assert response == mock_response def test_create_unauthorized_error(client, requests_mock): """Test that create raises UnauthorizedError for 401 responses""" - requests_mock.post(f"{client._base_url}/credentials", status_code=401, json={"error": "Unauthorized"}) + requests_mock.post( + f"{client._base_url}/credentials", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.create( - credential_name="azure1", credential_info={"api_type": "azure"}, credential_values={"api_key": "sk-123"} + credential_name="azure1", + credential_info={"api_type": "azure"}, + credential_values={"api_key": "sk-123"}, ) @@ -149,7 +173,11 @@ def test_delete_mock_response(client, requests_mock): def test_delete_unauthorized_error(client, requests_mock): """Test that delete raises UnauthorizedError for 401 responses""" - requests_mock.delete(f"{client._base_url}/credentials/azure1", status_code=401, json={"error": "Unauthorized"}) + requests_mock.delete( + f"{client._base_url}/credentials/azure1", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.delete(credential_name="azure1") @@ -170,11 +198,16 @@ def test_get_mock_response(client, requests_mock): mock_response = { "credential_name": "azure1", "credential_info": {"api_type": "azure"}, - "credential_values": {"api_key": "sk-123", "api_base": "https://example.azure.openai.com"}, + "credential_values": { + "api_key": "sk-123", + "api_base": "https://example.azure.openai.com", + }, "status": "active", } - requests_mock.get(f"{client._base_url}/credentials/by_name/azure1", json=mock_response) + requests_mock.get( + f"{client._base_url}/credentials/by_name/azure1", json=mock_response + ) response = client.get(credential_name="azure1") assert response == mock_response @@ -182,7 +215,11 @@ def test_get_mock_response(client, requests_mock): def test_get_unauthorized_error(client, requests_mock): """Test that get raises UnauthorizedError for 401 responses""" - requests_mock.get(f"{client._base_url}/credentials/by_name/azure1", status_code=401, json={"error": "Unauthorized"}) + requests_mock.get( + f"{client._base_url}/credentials/by_name/azure1", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.get(credential_name="azure1") diff --git a/tests/litellm/proxy/client/test_http_client.py b/tests/test_litellm/proxy/client/test_http_client.py similarity index 99% rename from tests/litellm/proxy/client/test_http_client.py rename to tests/test_litellm/proxy/client/test_http_client.py index 1a621959e03..30bb1489c58 100644 --- a/tests/litellm/proxy/client/test_http_client.py +++ b/tests/test_litellm/proxy/client/test_http_client.py @@ -1,9 +1,11 @@ """Tests for the HTTP client.""" import json + import pytest import requests import responses + from litellm.proxy.client.http_client import HTTPClient diff --git a/tests/litellm/proxy/client/test_http_commands.py b/tests/test_litellm/proxy/client/test_http_commands.py similarity index 98% rename from tests/litellm/proxy/client/test_http_commands.py rename to tests/test_litellm/proxy/client/test_http_commands.py index 8a2228bc5c2..f305efb3354 100644 --- a/tests/litellm/proxy/client/test_http_commands.py +++ b/tests/test_litellm/proxy/client/test_http_commands.py @@ -1,6 +1,7 @@ """Tests for the HTTP command group.""" import json + import pytest import responses from click.testing import CliRunner @@ -113,4 +114,4 @@ def test_request_invalid_header(runner): obj={"base_url": "http://localhost:4000", "api_key": "sk-test-key"}, ) assert result.exit_code == 2 # Click error code for invalid parameter - assert "Invalid header format" in result.output \ No newline at end of file + assert "Invalid header format" in result.output diff --git a/tests/litellm/proxy/client/test_keys.py b/tests/test_litellm/proxy/client/test_keys.py similarity index 88% rename from tests/litellm/proxy/client/test_keys.py rename to tests/test_litellm/proxy/client/test_keys.py index 85e4c371bb6..7ce4452bdc2 100644 --- a/tests/litellm/proxy/client/test_keys.py +++ b/tests/test_litellm/proxy/client/test_keys.py @@ -1,7 +1,8 @@ import pytest import requests -from litellm.proxy.client.keys import KeysManagementClient + from litellm.proxy.client.exceptions import UnauthorizedError +from litellm.proxy.client.keys import KeysManagementClient @pytest.fixture @@ -82,9 +83,14 @@ def test_list_request_filters(client): def test_list_request_flags(client): """Test list request with boolean flag parameters""" - request = client.list(return_full_object=True, include_team_keys=False, return_request=True) + request = client.list( + return_full_object=True, include_team_keys=False, return_request=True + ) - assert request.params == {"return_full_object": "true", "include_team_keys": "false"} + assert request.params == { + "return_full_object": "true", + "include_team_keys": "false", + } def test_list_request_all_parameters(client): @@ -170,7 +176,9 @@ def test_list_mock_response_filtered(client, requests_mock): requests_mock.get( f"{client._base_url}/key/list", json=mock_response, - additional_matcher=lambda r: (r.qs.get("user_id") == ["user123"] and r.qs.get("team_id") == ["team456"]), + additional_matcher=lambda r: ( + r.qs.get("user_id") == ["user123"] and r.qs.get("team_id") == ["team456"] + ), ) response = client.list(user_id="user123", team_id="team456") @@ -179,7 +187,9 @@ def test_list_mock_response_filtered(client, requests_mock): def test_list_unauthorized_error(client, requests_mock): """Test that list raises UnauthorizedError for 401 responses""" - requests_mock.get(f"{client._base_url}/key/list", status_code=401, json={"error": "Unauthorized"}) + requests_mock.get( + f"{client._base_url}/key/list", status_code=401, json={"error": "Unauthorized"} + ) with pytest.raises(UnauthorizedError): client.list() @@ -252,7 +262,11 @@ def test_generate_mock_response(client, requests_mock): def test_generate_unauthorized_error(client, requests_mock): """Test that generate raises UnauthorizedError for 401 responses""" - requests_mock.post(f"{client._base_url}/key/generate", status_code=401, json={"error": "Unauthorized"}) + requests_mock.post( + f"{client._base_url}/key/generate", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.generate() @@ -289,14 +303,20 @@ def test_delete_request_with_keys_and_aliases(client): """Test delete request with both keys and aliases""" keys_to_delete = ["key1", "key2"] aliases_to_delete = ["alias1", "alias2"] - request = client.delete(keys=keys_to_delete, key_aliases=aliases_to_delete, return_request=True) + request = client.delete( + keys=keys_to_delete, key_aliases=aliases_to_delete, return_request=True + ) assert request.json == {"keys": keys_to_delete, "key_aliases": aliases_to_delete} def test_delete_mock_response(client, requests_mock): """Test delete with a mocked successful response""" - mock_response = {"status": "success", "deleted_keys": ["key1", "key2"], "deleted_aliases": ["alias1"]} + mock_response = { + "status": "success", + "deleted_keys": ["key1", "key2"], + "deleted_aliases": ["alias1"], + } requests_mock.post(f"{client._base_url}/key/delete", json=mock_response) response = client.delete(keys=["key1", "key2"], key_aliases=["alias1"]) @@ -305,7 +325,11 @@ def test_delete_mock_response(client, requests_mock): def test_delete_unauthorized_error(client, requests_mock): """Test that delete raises UnauthorizedError for 401 responses""" - requests_mock.post(f"{client._base_url}/key/delete", status_code=401, json={"error": "Unauthorized"}) + requests_mock.post( + f"{client._base_url}/key/delete", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.delete(keys=["key-to-delete"]) @@ -336,7 +360,11 @@ def test_info_mock_response(client, requests_mock): def test_info_unauthorized_error(client, requests_mock): """Test that info raises UnauthorizedError for 401 responses""" - requests_mock.get(f"{client._base_url}/keys/info?key=test-key", status_code=401, json={"error": "Unauthorized"}) + requests_mock.get( + f"{client._base_url}/keys/info?key=test-key", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.info(key="test-key") @@ -344,7 +372,9 @@ def test_info_unauthorized_error(client, requests_mock): def test_info_server_error(client, requests_mock): """Test that info raises HTTPError for server errors""" requests_mock.get( - f"{client._base_url}/keys/info?key=test-key", status_code=500, json={"error": "Internal Server Error"} + f"{client._base_url}/keys/info?key=test-key", + status_code=500, + json={"error": "Internal Server Error"}, ) with pytest.raises(requests.exceptions.HTTPError): client.info(key="test-key") diff --git a/tests/litellm/proxy/client/test_model_groups.py b/tests/test_litellm/proxy/client/test_model_groups.py similarity index 88% rename from tests/litellm/proxy/client/test_model_groups.py rename to tests/test_litellm/proxy/client/test_model_groups.py index 11921b1a721..0afb253149b 100644 --- a/tests/litellm/proxy/client/test_model_groups.py +++ b/tests/test_litellm/proxy/client/test_model_groups.py @@ -1,5 +1,6 @@ import pytest import requests + from litellm.proxy.client import Client, ModelGroupsManagementClient from litellm.proxy.client.exceptions import UnauthorizedError @@ -51,7 +52,10 @@ def test_info_request_no_auth(base_url): "base_url,expected", [ ("http://localhost:8000", "http://localhost:8000/model_group/info"), - ("http://localhost:8000/", "http://localhost:8000/model_group/info"), # With trailing slash + ( + "http://localhost:8000/", + "http://localhost:8000/model_group/info", + ), # With trailing slash ("https://api.example.com", "https://api.example.com/model_group/info"), ("http://127.0.0.1:3000", "http://127.0.0.1:3000/model_group/info"), ], @@ -75,7 +79,10 @@ def test_info_with_mock_response(client, requests_mock): { "model_group_name": "azure-group", "models": ["azure-gpt-4", "azure-gpt-35"], - "litellm_params": {"api_base": "https://azure-endpoint.com", "api_version": "2023-05-15"}, + "litellm_params": { + "api_base": "https://azure-endpoint.com", + "api_version": "2023-05-15", + }, }, ] } @@ -90,7 +97,11 @@ def test_info_with_mock_response(client, requests_mock): def test_info_unauthorized_error(client, requests_mock): """Test that info raises UnauthorizedError for 401 responses""" - requests_mock.get(f"{client._base_url}/model_group/info", status_code=401, json={"error": "Invalid API key"}) + requests_mock.get( + f"{client._base_url}/model_group/info", + status_code=401, + json={"error": "Invalid API key"}, + ) with pytest.raises(UnauthorizedError) as exc_info: client.info() @@ -99,7 +110,11 @@ def test_info_unauthorized_error(client, requests_mock): def test_info_other_errors(client, requests_mock): """Test that info raises normal HTTPError for non-401 errors""" - requests_mock.get(f"{client._base_url}/model_group/info", status_code=500, json={"error": "Internal Server Error"}) + requests_mock.get( + f"{client._base_url}/model_group/info", + status_code=500, + json={"error": "Internal Server Error"}, + ) with pytest.raises(requests.exceptions.HTTPError) as exc_info: client.info() diff --git a/tests/litellm/proxy/client/test_models.py b/tests/test_litellm/proxy/client/test_models.py similarity index 81% rename from tests/litellm/proxy/client/test_models.py rename to tests/test_litellm/proxy/client/test_models.py index 0e1fbfe1ca3..afd042682b4 100644 --- a/tests/litellm/proxy/client/test_models.py +++ b/tests/test_litellm/proxy/client/test_models.py @@ -1,7 +1,8 @@ import pytest import requests + from litellm.proxy.client import Client, ModelsManagementClient -from litellm.proxy.client.exceptions import UnauthorizedError, NotFoundError +from litellm.proxy.client.exceptions import NotFoundError, UnauthorizedError @pytest.fixture @@ -51,7 +52,10 @@ def test_list_request_no_auth(base_url): "base_url,expected", [ ("http://localhost:8000", "http://localhost:8000/models"), - ("http://localhost:8000/", "http://localhost:8000/models"), # With trailing slash + ( + "http://localhost:8000/", + "http://localhost:8000/models", + ), # With trailing slash ("https://api.example.com", "https://api.example.com/models"), ("http://127.0.0.1:3000", "http://127.0.0.1:3000/models"), ], @@ -65,7 +69,12 @@ def test_list_url_variants(base_url, expected): def test_list_with_mock_response(client, requests_mock): """Test the full list execution with a mocked response""" - mock_data = {"data": [{"id": "gpt-4", "type": "model"}, {"id": "gpt-3.5-turbo", "type": "model"}]} + mock_data = { + "data": [ + {"id": "gpt-4", "type": "model"}, + {"id": "gpt-3.5-turbo", "type": "model"}, + ] + } requests_mock.get("http://localhost:8000/models", json=mock_data) response = client.list() @@ -76,7 +85,11 @@ def test_list_with_mock_response(client, requests_mock): def test_list_unauthorized_error(client, requests_mock): """Test that list raises UnauthorizedError for 401 responses""" - requests_mock.get("http://localhost:8000/models", status_code=401, json={"error": "Invalid API key"}) + requests_mock.get( + "http://localhost:8000/models", + status_code=401, + json={"error": "Invalid API key"}, + ) with pytest.raises(UnauthorizedError) as exc_info: client.list() @@ -85,7 +98,11 @@ def test_list_unauthorized_error(client, requests_mock): def test_list_other_errors(client, requests_mock): """Test that list raises normal HTTPError for non-401 errors""" - requests_mock.get("http://localhost:8000/models", status_code=500, json={"error": "Internal Server Error"}) + requests_mock.get( + "http://localhost:8000/models", + status_code=500, + json={"error": "Internal Server Error"}, + ) with pytest.raises(requests.exceptions.HTTPError) as exc_info: client.list() @@ -114,7 +131,12 @@ def test_client_initialization_strips_trailing_slash(): def test_list_with_mock_response(client, requests_mock): """Test the full list execution with a mocked response""" - mock_data = {"data": [{"id": "gpt-4", "type": "model"}, {"id": "gpt-3.5-turbo", "type": "model"}]} + mock_data = { + "data": [ + {"id": "gpt-4", "type": "model"}, + {"id": "gpt-3.5-turbo", "type": "model"}, + ] + } requests_mock.get("http://localhost:8000/models", json=mock_data) response = client.list() @@ -125,7 +147,11 @@ def test_list_with_mock_response(client, requests_mock): def test_list_unauthorized_error(client, requests_mock): """Test that list raises UnauthorizedError for 401 responses""" - requests_mock.get("http://localhost:8000/models", status_code=401, json={"error": "Invalid API key"}) + requests_mock.get( + "http://localhost:8000/models", + status_code=401, + json={"error": "Invalid API key"}, + ) with pytest.raises(UnauthorizedError) as exc_info: client.list() @@ -134,7 +160,11 @@ def test_list_unauthorized_error(client, requests_mock): def test_list_other_errors(client, requests_mock): """Test that list raises normal HTTPError for non-401 errors""" - requests_mock.get("http://localhost:8000/models", status_code=500, json={"error": "Internal Server Error"}) + requests_mock.get( + "http://localhost:8000/models", + status_code=500, + json={"error": "Internal Server Error"}, + ) with pytest.raises(requests.exceptions.HTTPError) as exc_info: client.list() @@ -192,7 +222,12 @@ def test_new_request_creation(client, base_url, api_key): model_params = {"model": "openai/gpt-4", "api_base": "https://api.openai.com/v1"} model_info = {"description": "GPT-4 model", "metadata": {"version": "1.0"}} - request = client.new(model_name=model_name, model_params=model_params, model_info=model_info, return_request=True) + request = client.new( + model_name=model_name, + model_params=model_params, + model_info=model_info, + return_request=True, + ) # Check request method and URL assert request.method == "POST" @@ -203,7 +238,11 @@ def test_new_request_creation(client, base_url, api_key): assert request.headers["Authorization"] == f"Bearer {api_key}" # Check request body - assert request.json == {"model_name": model_name, "litellm_params": model_params, "model_info": model_info} + assert request.json == { + "model_name": model_name, + "litellm_params": model_params, + "model_info": model_info, + } def test_new_without_model_info(client): @@ -211,7 +250,9 @@ def test_new_without_model_info(client): model_name = "gpt-4" model_params = {"model": "openai/gpt-4", "api_base": "https://api.openai.com/v1"} - request = client.new(model_name=model_name, model_params=model_params, return_request=True) + request = client.new( + model_name=model_name, model_params=model_params, return_request=True + ) # Check request body doesn't include model_info assert request.json == {"model_name": model_name, "litellm_params": model_params} @@ -237,7 +278,9 @@ def test_new_unauthorized_error(client, requests_mock): model_params = {"model": "openai/gpt-4"} # Mock a 401 response - requests_mock.post(f"{client._base_url}/model/new", status_code=401, json={"error": "Unauthorized"}) + requests_mock.post( + f"{client._base_url}/model/new", status_code=401, json={"error": "Unauthorized"} + ) with pytest.raises(UnauthorizedError): client.new(model_name=model_name, model_params=model_params) @@ -278,7 +321,11 @@ def test_delete_unauthorized_error(client, requests_mock): model_id = "model-123" # Mock a 401 response - requests_mock.post(f"{client._base_url}/model/delete", status_code=401, json={"error": "Unauthorized"}) + requests_mock.post( + f"{client._base_url}/model/delete", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.delete(model_id=model_id) @@ -289,7 +336,11 @@ def test_delete_404_error(client, requests_mock): model_id = "model-123" # Mock a 404 response - requests_mock.post(f"{client._base_url}/model/delete", status_code=404, json={"error": "Model not found"}) + requests_mock.post( + f"{client._base_url}/model/delete", + status_code=404, + json={"error": "Model not found"}, + ) with pytest.raises(NotFoundError) as exc_info: client.delete(model_id=model_id) @@ -317,7 +368,11 @@ def test_delete_other_errors(client, requests_mock): model_id = "model-123" # Mock a 500 response - requests_mock.post(f"{client._base_url}/model/delete", status_code=500, json={"error": "Internal Server Error"}) + requests_mock.post( + f"{client._base_url}/model/delete", + status_code=500, + json={"error": "Internal Server Error"}, + ) with pytest.raises(requests.exceptions.HTTPError) as exc_info: client.delete(model_id=model_id) @@ -343,7 +398,10 @@ def test_info_success(client, requests_mock): { "model_name": "gpt-4", "model_info": {"id": "model-123", "description": "GPT-4 model"}, - "litellm_params": {"model": "openai/gpt-4", "api_base": "https://api.openai.com/v1"}, + "litellm_params": { + "model": "openai/gpt-4", + "api_base": "https://api.openai.com/v1", + }, }, { "model_name": "gpt-3.5-turbo", @@ -364,7 +422,11 @@ def test_info_success(client, requests_mock): def test_info_unauthorized(client, requests_mock): """Test that info raises UnauthorizedError for unauthorized requests""" - requests_mock.get(f"{client._base_url}/v1/model/info", status_code=401, json={"error": "Unauthorized"}) + requests_mock.get( + f"{client._base_url}/v1/model/info", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError) as exc_info: client.info() @@ -373,7 +435,11 @@ def test_info_unauthorized(client, requests_mock): def test_info_server_error(client, requests_mock): """Test that info raises HTTPError for server errors""" - requests_mock.get(f"{client._base_url}/v1/model/info", status_code=500, json={"error": "Internal Server Error"}) + requests_mock.get( + f"{client._base_url}/v1/model/info", + status_code=500, + json={"error": "Internal Server Error"}, + ) with pytest.raises(requests.exceptions.HTTPError) as exc_info: client.info() @@ -410,12 +476,16 @@ def test_get_invalid_params(): # Test with no parameters with pytest.raises(ValueError) as exc_info: client.get() - assert "Exactly one of model_id or model_name must be provided" in str(exc_info.value) + assert "Exactly one of model_id or model_name must be provided" in str( + exc_info.value + ) # Test with both parameters with pytest.raises(ValueError) as exc_info: client.get(model_id="123", model_name="gpt-4") - assert "Exactly one of model_id or model_name must be provided" in str(exc_info.value) + assert "Exactly one of model_id or model_name must be provided" in str( + exc_info.value + ) def test_get_success_by_id(client, requests_mock): @@ -427,7 +497,10 @@ def test_get_success_by_id(client, requests_mock): { "model_name": "gpt-4", "model_info": {"id": model_id}, - "litellm_params": {"model": "openai/gpt-4", "api_base": "https://api.openai.com/v1"}, + "litellm_params": { + "model": "openai/gpt-4", + "api_base": "https://api.openai.com/v1", + }, }, ] } @@ -444,7 +517,11 @@ def test_get_success_by_name(client, requests_mock): model_name = "gpt-4" mock_models = { "data": [ - {"model_name": model_name, "model_info": {"id": "model-123"}, "litellm_params": {"model": "openai/gpt-4"}} + { + "model_name": model_name, + "model_info": {"id": "model-123"}, + "litellm_params": {"model": "openai/gpt-4"}, + } ] } @@ -461,7 +538,11 @@ def test_get_not_found(client, requests_mock): # Mock successful response but with no matching model requests_mock.get( f"{client._base_url}/v1/model/info", - json={"data": [{"model_name": "gpt-3.5-turbo", "model_info": {"id": "other-model"}}]}, + json={ + "data": [ + {"model_name": "gpt-3.5-turbo", "model_info": {"id": "other-model"}} + ] + }, ) with pytest.raises(NotFoundError) as exc_info: @@ -474,7 +555,11 @@ def test_get_unauthorized(client, requests_mock): """Test that get raises UnauthorizedError for unauthorized requests""" model_id = "model-123" - requests_mock.get(f"{client._base_url}/v1/model/info", status_code=401, json={"error": "Unauthorized"}) + requests_mock.get( + f"{client._base_url}/v1/model/info", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError) as exc_info: client.get(model_id=model_id) @@ -485,7 +570,11 @@ def test_get_server_error(client, requests_mock): """Test that get raises HTTPError for server errors""" model_id = "model-123" - requests_mock.get(f"{client._base_url}/v1/model/info", status_code=500, json={"error": "Internal Server Error"}) + requests_mock.get( + f"{client._base_url}/v1/model/info", + status_code=500, + json={"error": "Internal Server Error"}, + ) with pytest.raises(requests.exceptions.HTTPError) as exc_info: client.get(model_id=model_id) @@ -498,7 +587,12 @@ def test_update_request_creation(client, base_url, api_key): model_params = {"model": "openai/gpt-4", "api_base": "https://api.openai.com/v1"} model_info = {"description": "Updated GPT-4 model", "metadata": {"version": "2.0"}} - request = client.update(model_id=model_id, model_params=model_params, model_info=model_info, return_request=True) + request = client.update( + model_id=model_id, + model_params=model_params, + model_info=model_info, + return_request=True, + ) # Check request method and URL assert request.method == "POST" @@ -509,7 +603,11 @@ def test_update_request_creation(client, base_url, api_key): assert request.headers["Authorization"] == f"Bearer {api_key}" # Check request body - assert request.json == {"id": model_id, "litellm_params": model_params, "model_info": model_info} + assert request.json == { + "id": model_id, + "litellm_params": model_params, + "model_info": model_info, + } def test_update_without_model_info(client): @@ -517,7 +615,9 @@ def test_update_without_model_info(client): model_id = "model-123" model_params = {"model": "openai/gpt-4", "api_base": "https://api.openai.com/v1"} - request = client.update(model_id=model_id, model_params=model_params, return_request=True) + request = client.update( + model_id=model_id, model_params=model_params, return_request=True + ) # Check request body doesn't include model_info assert request.json == {"id": model_id, "litellm_params": model_params} @@ -527,7 +627,11 @@ def test_update_mock_response(client, requests_mock): """Test update with a mocked successful response""" model_id = "model-123" model_params = {"model": "openai/gpt-4"} - mock_response = {"id": model_id, "status": "success", "message": "Model updated successfully"} + mock_response = { + "id": model_id, + "status": "success", + "message": "Model updated successfully", + } # Mock the POST request requests_mock.post(f"{client._base_url}/model/update", json=mock_response) @@ -543,7 +647,11 @@ def test_update_unauthorized_error(client, requests_mock): model_params = {"model": "openai/gpt-4"} # Mock a 401 response - requests_mock.post(f"{client._base_url}/model/update", status_code=401, json={"error": "Unauthorized"}) + requests_mock.post( + f"{client._base_url}/model/update", + status_code=401, + json={"error": "Unauthorized"}, + ) with pytest.raises(UnauthorizedError): client.update(model_id=model_id, model_params=model_params) @@ -555,7 +663,11 @@ def test_update_404_error(client, requests_mock): model_params = {"model": "openai/gpt-4"} # Mock a 404 response - requests_mock.post(f"{client._base_url}/model/update", status_code=404, json={"error": "Model not found"}) + requests_mock.post( + f"{client._base_url}/model/update", + status_code=404, + json={"error": "Model not found"}, + ) with pytest.raises(NotFoundError) as exc_info: client.update(model_id=model_id, model_params=model_params) @@ -585,7 +697,11 @@ def test_update_other_errors(client, requests_mock): model_params = {"model": "openai/gpt-4"} # Mock a 500 response - requests_mock.post(f"{client._base_url}/model/update", status_code=500, json={"error": "Internal Server Error"}) + requests_mock.post( + f"{client._base_url}/model/update", + status_code=500, + json={"error": "Internal Server Error"}, + ) with pytest.raises(requests.exceptions.HTTPError) as exc_info: client.update(model_id=model_id, model_params=model_params) diff --git a/tests/litellm/proxy/client/test_users.py b/tests/test_litellm/proxy/client/test_users.py similarity index 91% rename from tests/litellm/proxy/client/test_users.py rename to tests/test_litellm/proxy/client/test_users.py index 8cc57713430..5bcd783b914 100644 --- a/tests/litellm/proxy/client/test_users.py +++ b/tests/test_litellm/proxy/client/test_users.py @@ -1,11 +1,19 @@ +from unittest.mock import MagicMock, patch + import pytest -from unittest.mock import patch, MagicMock -from litellm.proxy.client.users import UsersManagementClient, UnauthorizedError, NotFoundError + +from litellm.proxy.client.users import ( + NotFoundError, + UnauthorizedError, + UsersManagementClient, +) + @pytest.fixture def client(): return UsersManagementClient(base_url="http://localhost:4000", api_key="sk-test") + @patch("requests.get") def test_list_users_success(mock_get, client): mock_get.return_value.status_code = 200 @@ -14,6 +22,7 @@ def test_list_users_success(mock_get, client): assert users == [{"user_id": "u1"}] mock_get.assert_called_once() + @patch("requests.get") def test_list_users_unauthorized(mock_get, client): mock_get.return_value.status_code = 401 @@ -21,6 +30,7 @@ def test_list_users_unauthorized(mock_get, client): with pytest.raises(UnauthorizedError): client.list_users() + @patch("requests.get") def test_get_user_success(mock_get, client): mock_get.return_value.status_code = 200 @@ -29,6 +39,7 @@ def test_get_user_success(mock_get, client): assert user["user_id"] == "u1" mock_get.assert_called_once() + @patch("requests.get") def test_get_user_404(mock_get, client): mock_get.return_value.status_code = 404 @@ -36,6 +47,7 @@ def test_get_user_404(mock_get, client): with pytest.raises(NotFoundError): client.get_user(user_id="u1") + @patch("requests.post") def test_create_user_success(mock_post, client): mock_post.return_value.status_code = 200 @@ -44,6 +56,7 @@ def test_create_user_success(mock_post, client): assert user["user_id"] == "u1" mock_post.assert_called_once() + @patch("requests.post") def test_create_user_unauthorized(mock_post, client): mock_post.return_value.status_code = 401 @@ -51,6 +64,7 @@ def test_create_user_unauthorized(mock_post, client): with pytest.raises(UnauthorizedError): client.create_user({"user_email": "a@b.com"}) + @patch("requests.post") def test_delete_user_success(mock_post, client): mock_post.return_value.status_code = 200 @@ -59,9 +73,10 @@ def test_delete_user_success(mock_post, client): assert result["deleted"] == 1 mock_post.assert_called_once() + @patch("requests.post") def test_delete_user_unauthorized(mock_post, client): mock_post.return_value.status_code = 401 mock_post.return_value.text = "unauthorized" with pytest.raises(UnauthorizedError): - client.delete_user(["u1"]) \ No newline at end of file + client.delete_user(["u1"]) diff --git a/tests/litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py similarity index 100% rename from tests/litellm/proxy/common_utils/test_http_parsing_utils.py rename to tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py diff --git a/tests/litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py similarity index 99% rename from tests/litellm/proxy/common_utils/test_reset_budget_job.py rename to tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 9f25279f1dd..fbf02d66127 100644 --- a/tests/litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -198,4 +198,4 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client): # Check that all spends were reset to 0 assert mock_prisma_client.updated_data["key"][0].spend == 0.0 assert mock_prisma_client.updated_data["user"][0].spend == 0.0 - assert mock_prisma_client.updated_data["team"][0].spend == 0.0 \ No newline at end of file + assert mock_prisma_client.updated_data["team"][0].spend == 0.0 diff --git a/tests/litellm/proxy/common_utils/test_timezone_utils.py b/tests/test_litellm/proxy/common_utils/test_timezone_utils.py similarity index 100% rename from tests/litellm/proxy/common_utils/test_timezone_utils.py rename to tests/test_litellm/proxy/common_utils/test_timezone_utils.py diff --git a/tests/litellm/proxy/db/db_transaction_queue/test_base_update_queue.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_base_update_queue.py similarity index 100% rename from tests/litellm/proxy/db/db_transaction_queue/test_base_update_queue.py rename to tests/test_litellm/proxy/db/db_transaction_queue/test_base_update_queue.py diff --git a/tests/litellm/proxy/db/db_transaction_queue/test_daily_spend_update_queue.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_daily_spend_update_queue.py similarity index 100% rename from tests/litellm/proxy/db/db_transaction_queue/test_daily_spend_update_queue.py rename to tests/test_litellm/proxy/db/db_transaction_queue/test_daily_spend_update_queue.py diff --git a/tests/litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py similarity index 100% rename from tests/litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py rename to tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py diff --git a/tests/litellm/proxy/db/db_transaction_queue/test_spend_update_queue.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_spend_update_queue.py similarity index 100% rename from tests/litellm/proxy/db/db_transaction_queue/test_spend_update_queue.py rename to tests/test_litellm/proxy/db/db_transaction_queue/test_spend_update_queue.py diff --git a/tests/litellm/proxy/db/test_check_migration.py b/tests/test_litellm/proxy/db/test_check_migration.py similarity index 100% rename from tests/litellm/proxy/db/test_check_migration.py rename to tests/test_litellm/proxy/db/test_check_migration.py diff --git a/tests/litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py similarity index 100% rename from tests/litellm/proxy/db/test_db_spend_update_writer.py rename to tests/test_litellm/proxy/db/test_db_spend_update_writer.py diff --git a/tests/litellm/proxy/db/test_exception_handler.py b/tests/test_litellm/proxy/db/test_exception_handler.py similarity index 100% rename from tests/litellm/proxy/db/test_exception_handler.py rename to tests/test_litellm/proxy/db/test_exception_handler.py diff --git a/tests/litellm/proxy/db/test_prisma_client.py b/tests/test_litellm/proxy/db/test_prisma_client.py similarity index 100% rename from tests/litellm/proxy/db/test_prisma_client.py rename to tests/test_litellm/proxy/db/test_prisma_client.py diff --git a/tests/litellm/proxy/experimental/mcp_server/test_tool_registry.py b/tests/test_litellm/proxy/experimental/mcp_server/test_tool_registry.py similarity index 100% rename from tests/litellm/proxy/experimental/mcp_server/test_tool_registry.py rename to tests/test_litellm/proxy/experimental/mcp_server/test_tool_registry.py diff --git a/tests/litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py similarity index 100% rename from tests/litellm/proxy/guardrails/test_guardrail_endpoints.py rename to tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py diff --git a/tests/litellm/proxy/guardrails/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py similarity index 100% rename from tests/litellm/proxy/guardrails/test_init_guardrails.py rename to tests/test_litellm/proxy/guardrails/test_init_guardrails.py diff --git a/tests/litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py similarity index 99% rename from tests/litellm/proxy/health_endpoints/test_health_endpoints.py rename to tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index e2dd429357a..2502b4e34e6 100644 --- a/tests/litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -50,7 +50,6 @@ async def test_db_health_readiness_check_with_prisma_error(prisma_error): "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": True}, ): - # Call the function result = await _db_health_readiness_check() @@ -92,7 +91,6 @@ async def test_db_health_readiness_check_with_error_and_flag_off(prisma_error): "litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": False}, ): - # The function should raise the exception with pytest.raises(Exception) as excinfo: await _db_health_readiness_check() diff --git a/tests/litellm/proxy/hooks/test_parallel_request_limiter_v2.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v2.py similarity index 100% rename from tests/litellm/proxy/hooks/test_parallel_request_limiter_v2.py rename to tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v2.py diff --git a/tests/litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py similarity index 100% rename from tests/litellm/proxy/hooks/test_proxy_track_cost_callback.py rename to tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py diff --git a/tests/litellm/proxy/management_endpoints/scim/test_scim_transformations.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py similarity index 100% rename from tests/litellm/proxy/management_endpoints/scim/test_scim_transformations.py rename to tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py diff --git a/tests/litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py similarity index 100% rename from tests/litellm/proxy/management_endpoints/test_common_daily_activity.py rename to tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py diff --git a/tests/litellm/proxy/management_endpoints/test_customer_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py similarity index 100% rename from tests/litellm/proxy/management_endpoints/test_customer_endpoints.py rename to tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py diff --git a/tests/litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py similarity index 100% rename from tests/litellm/proxy/management_endpoints/test_internal_user_endpoints.py rename to tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py diff --git a/tests/litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py similarity index 100% rename from tests/litellm/proxy/management_endpoints/test_key_management_endpoints.py rename to tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py diff --git a/tests/litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py similarity index 96% rename from tests/litellm/proxy/management_endpoints/test_model_management_endpoints.py rename to tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 129d0d205a5..e9bb69abd0e 100644 --- a/tests/litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -55,10 +55,10 @@ class MockLLMRouter: self.model_list = ["model1", "model2"] self.model_names = {"model1": True, "model2": True} self.cleared = False - + def get_deployment(self, model_id): return {"model_id": model_id} if model_id in self.model_list else None - + def delete_deployment(self, id): if id in self.model_list: self.model_list.remove(id) @@ -69,7 +69,7 @@ class MockProxyConfig: def __init__(self, success=True): self.success = success self.deployment_called = False - + async def add_deployment(self, prisma_client, proxy_logging_obj): self.deployment_called = True if not self.success: @@ -361,11 +361,12 @@ class TestDeleteTeamModelAlias: mock_db = mock_prisma.db.litellm_modeltable assert len(mock_db.update_calls) == 0 + class TestClearCache: """ Tests for the clear_cache function in model_management_endpoints.py """ - + @pytest.mark.asyncio async def test_clear_cache_success(self): """ @@ -373,25 +374,24 @@ class TestClearCache: """ mock_router = MagicMock() mock_router.model_list = ["openai/gpt-4o", "openai/gpt-4o-mini"] - + mock_config = MagicMock() mock_config.add_deployment = AsyncMock(return_value=True) - + mock_prisma = MagicMock() mock_logging = MagicMock() - - with patch("litellm.proxy.proxy_server.llm_router", mock_router), \ - patch("litellm.proxy.proxy_server.proxy_config", mock_config), \ - patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), \ - patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_logging), \ - patch("litellm.proxy.proxy_server.verbose_proxy_logger"): - + + with patch("litellm.proxy.proxy_server.llm_router", mock_router), patch( + "litellm.proxy.proxy_server.proxy_config", mock_config + ), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", mock_logging + ), patch( + "litellm.proxy.proxy_server.verbose_proxy_logger" + ): await clear_cache() - - + assert len(mock_router.model_list) == 0 - + mock_config.add_deployment.assert_called_once_with( - prisma_client=mock_prisma, - proxy_logging_obj=mock_logging + prisma_client=mock_prisma, proxy_logging_obj=mock_logging ) diff --git a/tests/litellm/proxy/management_endpoints/test_tag_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py similarity index 100% rename from tests/litellm/proxy/management_endpoints/test_tag_management_endpoints.py rename to tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py diff --git a/tests/litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py similarity index 100% rename from tests/litellm/proxy/management_endpoints/test_team_endpoints.py rename to tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py diff --git a/tests/litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py similarity index 100% rename from tests/litellm/proxy/management_endpoints/test_ui_sso.py rename to tests/test_litellm/proxy/management_endpoints/test_ui_sso.py diff --git a/tests/litellm/proxy/middleware/test_prometheus_auth_middleware.py b/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py similarity index 100% rename from tests/litellm/proxy/middleware/test_prometheus_auth_middleware.py rename to tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py diff --git a/tests/litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py similarity index 100% rename from tests/litellm/proxy/openai_files_endpoint/test_files_endpoint.py rename to tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py diff --git a/tests/litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py similarity index 100% rename from tests/litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py rename to tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py diff --git a/tests/litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py similarity index 100% rename from tests/litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py rename to tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py diff --git a/tests/litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py similarity index 100% rename from tests/litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py rename to tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py diff --git a/tests/litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py similarity index 100% rename from tests/litellm/proxy/spend_tracking/test_spend_management_endpoints.py rename to tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py diff --git a/tests/litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py similarity index 100% rename from tests/litellm/proxy/spend_tracking/test_spend_tracking_utils.py rename to tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py diff --git a/tests/litellm/proxy/test_caching_routes.py b/tests/test_litellm/proxy/test_caching_routes.py similarity index 100% rename from tests/litellm/proxy/test_caching_routes.py rename to tests/test_litellm/proxy/test_caching_routes.py diff --git a/tests/litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py similarity index 93% rename from tests/litellm/proxy/test_common_request_processing.py rename to tests/test_litellm/proxy/test_common_request_processing.py index 299bdda917b..e46e0c28b13 100644 --- a/tests/litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -1,13 +1,13 @@ import copy import uuid -import pytest -import litellm from unittest.mock import AsyncMock, MagicMock -from fastapi import Request -from litellm.integrations.opentelemetry import UserAPIKeyAuth -from fastapi import status +import pytest +from fastapi import Request, status from fastapi.responses import StreamingResponse + +import litellm +from litellm.integrations.opentelemetry import UserAPIKeyAuth from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, ProxyConfig, @@ -18,12 +18,10 @@ from litellm.proxy.utils import ProxyLogging class TestProxyBaseLLMRequestProcessing: - @pytest.mark.asyncio async def test_common_processing_pre_call_logic_pre_call_hook_receives_litellm_call_id( self, monkeypatch ): - processing_obj = ProxyBaseLLMRequestProcessing(data={}) mock_request = MagicMock(spec=Request) mock_request.headers = {} @@ -52,15 +50,16 @@ class TestProxyBaseLLMRequestProcessing: route_type = "acompletion" # Call the actual method. - returned_data, logging_obj = ( - await processing_obj.common_processing_pre_call_logic( - request=mock_request, - general_settings=mock_general_settings, - user_api_key_dict=mock_user_api_key_dict, - proxy_logging_obj=mock_proxy_logging_obj, - proxy_config=mock_proxy_config, - route_type=route_type, - ) + ( + returned_data, + logging_obj, + ) = await processing_obj.common_processing_pre_call_logic( + request=mock_request, + general_settings=mock_general_settings, + user_api_key_dict=mock_user_api_key_dict, + proxy_logging_obj=mock_proxy_logging_obj, + proxy_config=mock_proxy_config, + route_type=route_type, ) mock_proxy_logging_obj.pre_call_hook.assert_called_once() @@ -125,7 +124,7 @@ class TestCommonRequestProcessingHelpers: ( 'data: {"error": {"code": null, "message": "code is null"}}', None, - ), # Error with null code + ), # Error with null code ], ) async def test_parse_event_data_for_error(self, event_line, expected_code): @@ -207,10 +206,10 @@ class TestCommonRequestProcessingHelpers: assert len(content) == 2 # Use json.dumps to match the formatting in create_streaming_response's exception handler import json + assert content[0] == f"data: {json.dumps(expected_error_data)}\n\n" assert content[1] == "data: [DONE]\n\n" - async def test_create_streaming_response_first_chunk_error_string_code(self): async def mock_generator(): yield 'data: {"error": {"code": "429", "message": "too many requests"}}\n\n' @@ -262,7 +261,7 @@ class TestCommonRequestProcessingHelpers: response = await create_streaming_response( mock_generator(), "text/event-stream", {} ) - assert response.status_code == status.HTTP_200_OK # Default status + assert response.status_code == status.HTTP_200_OK # Default status content = await self.consume_stream(response) assert content == ["data: [DONE]\n\n"] @@ -275,7 +274,7 @@ class TestCommonRequestProcessingHelpers: response = await create_streaming_response( mock_generator(), "text/event-stream", {} ) - assert response.status_code == status.HTTP_200_OK # Default status + assert response.status_code == status.HTTP_200_OK # Default status content = await self.consume_stream(response) assert content == [ "data: \n\n", diff --git a/tests/litellm/proxy/test_configs/test_config_no_auth.yaml b/tests/test_litellm/proxy/test_configs/test_config_no_auth.yaml similarity index 100% rename from tests/litellm/proxy/test_configs/test_config_no_auth.yaml rename to tests/test_litellm/proxy/test_configs/test_config_no_auth.yaml diff --git a/tests/litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py similarity index 95% rename from tests/litellm/proxy/test_litellm_pre_call_utils.py rename to tests/test_litellm/proxy/test_litellm_pre_call_utils.py index a390c2e072a..19dadf17b0e 100644 --- a/tests/litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -1,18 +1,18 @@ +import asyncio +import copy import json import os import sys from unittest.mock import MagicMock, patch import pytest -import asyncio -import copy from fastapi import Request from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import ( _get_enforced_params, - check_if_token_is_service_account, add_litellm_data_to_request, + check_if_token_is_service_account, ) sys.path.insert( @@ -168,24 +168,21 @@ async def test_add_litellm_data_to_request_audio_transcription_multipart(): request_mock.query_params = {} request_mock.headers = { "Content-Type": "multipart/form-data", - "Authorization": "Bearer sk-1234" + "Authorization": "Bearer sk-1234", } request_mock.client = MagicMock() request_mock.client.host = "127.0.0.1" # Simulate multipart data (metadata as string) metadata_dict = { - "tags": [ - "jobID:214590dsff09fds", - "taskName:run_page_classification" - ] + "tags": ["jobID:214590dsff09fds", "taskName:run_page_classification"] } stringified_metadata = json.dumps(metadata_dict) data = { "model": "fake-openai-endpoint", "metadata": stringified_metadata, # Simulating multipart-form field - "file": b"Fake audio bytes" + "file": b"Fake audio bytes", } user_api_key_dict = UserAPIKeyAuth( @@ -214,4 +211,7 @@ async def test_add_litellm_data_to_request_audio_transcription_multipart(): assert isinstance(metadata_field, dict) assert "tags" in metadata_field - assert metadata_field["tags"] == ["jobID:214590dsff09fds", "taskName:run_page_classification"] + assert metadata_field["tags"] == [ + "jobID:214590dsff09fds", + "taskName:run_page_classification", + ] diff --git a/tests/litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py similarity index 91% rename from tests/litellm/proxy/test_proxy_cli.py rename to tests/test_litellm/proxy/test_proxy_cli.py index 17be8d75c23..70a39df485b 100644 --- a/tests/litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -1,6 +1,7 @@ import os import sys from unittest.mock import MagicMock, patch + import pytest sys.path.insert( @@ -11,7 +12,6 @@ from litellm.proxy.proxy_cli import ProxyInitializationHelpers class TestProxyInitializationHelpers: - @patch("importlib.metadata.version") @patch("click.echo") def test_echo_litellm_version(self, mock_echo, mock_version): @@ -156,7 +156,6 @@ class TestProxyInitializationHelpers: with patch("sys.platform", "linux"): assert ProxyInitializationHelpers._get_loop_type() == "uvloop" - @patch.dict(os.environ, {}, clear=True) def test_database_url_construction_with_special_characters(self): # Setup environment variables with special characters that need escaping @@ -164,15 +163,16 @@ class TestProxyInitializationHelpers: "DATABASE_HOST": "localhost:5432", "DATABASE_USERNAME": "user@with+special", "DATABASE_PASSWORD": "pass&word!@#$%", - "DATABASE_NAME": "db_name/test" + "DATABASE_NAME": "db_name/test", } with patch.dict(os.environ, test_env): # Call the relevant function - we'll need to extract the database URL construction logic # This is simulating what happens in the run_server function when database_url is None - from litellm.proxy.proxy_cli import append_query_params import urllib.parse + from litellm.proxy.proxy_cli import append_query_params + database_host = os.environ["DATABASE_HOST"] database_username = os.environ["DATABASE_USERNAME"] database_password = os.environ["DATABASE_PASSWORD"] @@ -191,10 +191,7 @@ class TestProxyInitializationHelpers: assert database_url == expected_url # Test appending query parameters - params = { - "connection_limit": 10, - "pool_timeout": 60 - } + params = {"connection_limit": 10, "pool_timeout": 60} modified_url = append_query_params(database_url, params) assert "connection_limit=10" in modified_url assert "pool_timeout=60" in modified_url @@ -204,42 +201,47 @@ class TestProxyInitializationHelpers: def test_skip_server_startup(self, mock_print, mock_uvicorn_run): """Test that the skip_server_startup flag prevents server startup when True""" from click.testing import CliRunner + from litellm.proxy.proxy_cli import run_server - + runner = CliRunner() mock_app = MagicMock() mock_proxy_config = MagicMock() mock_key_mgmt = MagicMock() mock_save_worker_config = MagicMock() - - with patch.dict('sys.modules', { - 'proxy_server': MagicMock( + + with patch.dict( + "sys.modules", + { + "proxy_server": MagicMock( app=mock_app, ProxyConfig=mock_proxy_config, KeyManagementSettings=mock_key_mgmt, - save_worker_config=mock_save_worker_config + save_worker_config=mock_save_worker_config, ) - }), \ - patch("litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args") as mock_get_args: - + }, + ), patch( + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + ) as mock_get_args: mock_get_args.return_value = { - "app": "litellm.proxy.proxy_server:app", - "host": "localhost", - "port": 8000 + "app": "litellm.proxy.proxy_server:app", + "host": "localhost", + "port": 8000, } - + result = runner.invoke(run_server, ["--local", "--skip_server_startup"]) - + assert result.exit_code == 0 mock_uvicorn_run.assert_not_called() - mock_print.assert_any_call("LiteLLM: Setup complete. Skipping server startup as requested.") - + mock_print.assert_any_call( + "LiteLLM: Setup complete. Skipping server startup as requested." + ) + mock_uvicorn_run.reset_mock() mock_print.reset_mock() - + result = runner.invoke(run_server, ["--local"]) - + assert result.exit_code == 0 mock_uvicorn_run.assert_called_once() - diff --git a/tests/litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py similarity index 99% rename from tests/litellm/proxy/test_proxy_server.py rename to tests/test_litellm/proxy/test_proxy_server.py index ba787bc1e1e..21758bbc8b4 100644 --- a/tests/litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -48,12 +48,14 @@ example_embedding_result = { "usage": {"prompt_tokens": 5, "total_tokens": 5}, } + def mock_patch_aembedding(): return mock.patch( "litellm.proxy.proxy_server.llm_router.aembedding", return_value=example_embedding_result, ) + @pytest.fixture(scope="function") def client_no_auth(): # Assuming litellm.proxy.proxy_server is an object diff --git a/tests/litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py similarity index 87% rename from tests/litellm/proxy/test_route_llm_request.py rename to tests/test_litellm/proxy/test_route_llm_request.py index db6a23f9253..c99815bcd91 100644 --- a/tests/litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -45,32 +45,26 @@ async def test_route_request_dynamic_credentials(route_type): # Now assert that the dynamic method was called once with the expected kwargs. getattr(llm_router, route_type).assert_called_once_with(**data) + @pytest.mark.asyncio async def test_route_request_no_model_required(): """Test route types that don't require model parameter""" - test_cases = [ - "amoderation", - "aget_responses", - "adelete_responses" - ] - + test_cases = ["amoderation", "aget_responses", "adelete_responses"] + for route_type in test_cases: # Test data without model parameter - data = { - "input": "test input", - "api_key": "test-key" - } - + data = {"input": "test input", "api_key": "test-key"} + llm_router = MagicMock() getattr(llm_router, route_type).return_value = "fake_response" - + response = await route_request(data, llm_router, None, route_type) - + # Verify response assert response == "fake_response" # Verify the method was called with correct parameters getattr(llm_router, route_type).assert_called_once_with(**data) - + # Reset mock for next iteration llm_router.reset_mock() @@ -78,17 +72,13 @@ async def test_route_request_no_model_required(): @pytest.mark.asyncio async def test_route_request_no_model_required_with_router_settings(): """Test route types that don't require model parameter with router settings""" - test_cases = [ - "amoderation", - "aget_responses", - "adelete_responses" - ] + test_cases = ["amoderation", "aget_responses", "adelete_responses"] for route_type in test_cases: # Test data with model parameter (it will be ignored for these route types) data = { "input": "test input", - "model": "test-model" # Include dummy model to avoid KeyError + "model": "test-model", # Include dummy model to avoid KeyError } llm_router = MagicMock() @@ -111,4 +101,4 @@ async def test_route_request_no_model_required_with_router_settings(): getattr(llm_router, route_type).assert_called_once_with(**data) # Reset the mock for the next route - llm_router.reset_mock() \ No newline at end of file + llm_router.reset_mock() diff --git a/tests/litellm/proxy/test_spend_log_cleanup.py b/tests/test_litellm/proxy/test_spend_log_cleanup.py similarity index 85% rename from tests/litellm/proxy/test_spend_log_cleanup.py rename to tests/test_litellm/proxy/test_spend_log_cleanup.py index 0b97f7c2c86..3c66d9baef1 100644 --- a/tests/litellm/proxy/test_spend_log_cleanup.py +++ b/tests/test_litellm/proxy/test_spend_log_cleanup.py @@ -2,10 +2,13 @@ Test cases for spend log cleanup functionality """ +from datetime import UTC, datetime, timedelta, timezone +from unittest.mock import AsyncMock, MagicMock + import pytest -from datetime import datetime, timedelta, UTC, timezone + from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup -from unittest.mock import MagicMock, AsyncMock + @pytest.mark.asyncio async def test_should_delete_spend_logs(): @@ -14,26 +17,34 @@ async def test_should_delete_spend_logs(): assert cleaner._should_delete_spend_logs() is False # Test case 2: Valid seconds string - cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "3600s"}) + cleaner = SpendLogCleanup( + general_settings={"maximum_spend_logs_retention_period": "3600s"} + ) assert cleaner._should_delete_spend_logs() is True # Test case 3: Valid days string - cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "30d"}) + cleaner = SpendLogCleanup( + general_settings={"maximum_spend_logs_retention_period": "30d"} + ) assert cleaner._should_delete_spend_logs() is True # Test case 4: Valid hours string - cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "24h"}) + cleaner = SpendLogCleanup( + general_settings={"maximum_spend_logs_retention_period": "24h"} + ) assert cleaner._should_delete_spend_logs() is True # Test case 5: Invalid format - cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "invalid"}) + cleaner = SpendLogCleanup( + general_settings={"maximum_spend_logs_retention_period": "invalid"} + ) assert cleaner._should_delete_spend_logs() is False @pytest.mark.asyncio async def test_cleanup_old_spend_logs_batch_deletion(): from types import SimpleNamespace - from unittest.mock import MagicMock, AsyncMock, patch + from unittest.mock import AsyncMock, MagicMock, patch # Setup Prisma client mock_prisma_client = MagicMock() @@ -49,7 +60,7 @@ async def test_cleanup_old_spend_logs_batch_deletion(): mock_spendlogs.find_many.side_effect = [ mock_logs[:1000], # Batch 1 mock_logs[1000:], # Batch 2 - [] # Done + [], # Done ] # Wire up mocks @@ -80,6 +91,7 @@ async def test_cleanup_old_spend_logs_batch_deletion(): where={"request_id": {"in": [f"req_{i}" for i in range(1000, 1500)]}} ) + @pytest.mark.asyncio async def test_cleanup_old_spend_logs_retention_period_cutoff(): """ @@ -111,7 +123,10 @@ async def test_cleanup_old_spend_logs_retention_period_cutoff(): # Verify the cutoff date is correct cutoff_date = mock_spendlogs.find_many.call_args[1]["where"]["startTime"]["lt"] expected_cutoff = datetime.now(timezone.utc) - timedelta(seconds=86400) - assert abs((cutoff_date - expected_cutoff).total_seconds()) < 1 # Allow 1 second difference for test execution time + assert ( + abs((cutoff_date - expected_cutoff).total_seconds()) < 1 + ) # Allow 1 second difference for test execution time + @pytest.mark.asyncio async def test_cleanup_old_spend_logs_no_retention_period(): @@ -126,4 +141,4 @@ async def test_cleanup_old_spend_logs_no_retention_period(): await cleaner.cleanup_old_spend_logs(mock_prisma_client) mock_prisma_client.db.litellm_spendlogs.find_many.assert_not_called() - mock_prisma_client.db.litellm_spendlogs.delete.assert_not_called() \ No newline at end of file + mock_prisma_client.db.litellm_spendlogs.delete.assert_not_called() diff --git a/tests/litellm/proxy/test_team_member_update.py b/tests/test_litellm/proxy/test_team_member_update.py similarity index 92% rename from tests/litellm/proxy/test_team_member_update.py rename to tests/test_litellm/proxy/test_team_member_update.py index 7f534b0b618..77ca01beeeb 100644 --- a/tests/litellm/proxy/test_team_member_update.py +++ b/tests/test_litellm/proxy/test_team_member_update.py @@ -2,15 +2,16 @@ import pytest from fastapi import HTTPException from starlette.requests import Request -from litellm.proxy.management_endpoints.team_endpoints import team_member_update -from litellm.proxy._types import TeamMemberUpdateRequest import litellm.proxy.proxy_server as proxy_server +from litellm.proxy._types import TeamMemberUpdateRequest +from litellm.proxy.management_endpoints.team_endpoints import team_member_update + @pytest.mark.asyncio async def test_team_member_update_admin_requires_premium(monkeypatch): # Arrange: patch prisma_client and premium_user - monkeypatch.setattr(proxy_server, 'prisma_client', object()) - monkeypatch.setattr(proxy_server, 'premium_user', False) + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr(proxy_server, "premium_user", False) # Create a request body that tries to set role=admin data = TeamMemberUpdateRequest( diff --git a/tests/litellm/proxy/types_utils/test_litellm_proxy_types_utils.py b/tests/test_litellm/proxy/types_utils/test_litellm_proxy_types_utils.py similarity index 100% rename from tests/litellm/proxy/types_utils/test_litellm_proxy_types_utils.py rename to tests/test_litellm/proxy/types_utils/test_litellm_proxy_types_utils.py diff --git a/tests/litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py similarity index 99% rename from tests/litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py rename to tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 0b29ef6c7a6..543258640ed 100644 --- a/tests/litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -77,7 +77,6 @@ def mock_auth(monkeypatch): class TestProxySettingEndpoints: - def test_get_internal_user_settings(self, mock_proxy_config, mock_auth): """Test getting the internal user settings""" response = client.get("/get/internal_user_settings") diff --git a/tests/litellm/readme.md b/tests/test_litellm/readme.md similarity index 100% rename from tests/litellm/readme.md rename to tests/test_litellm/readme.md diff --git a/tests/litellm/responses/test_responses_utils.py b/tests/test_litellm/responses/test_responses_utils.py similarity index 100% rename from tests/litellm/responses/test_responses_utils.py rename to tests/test_litellm/responses/test_responses_utils.py diff --git a/tests/litellm/router_strategy/test_base_routing_strategy.py b/tests/test_litellm/router_strategy/test_base_routing_strategy.py similarity index 100% rename from tests/litellm/router_strategy/test_base_routing_strategy.py rename to tests/test_litellm/router_strategy/test_base_routing_strategy.py diff --git a/tests/litellm/router_utils/pre_call_checks/test_responses_api_deployment_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_responses_api_deployment_check.py similarity index 100% rename from tests/litellm/router_utils/pre_call_checks/test_responses_api_deployment_check.py rename to tests/test_litellm/router_utils/pre_call_checks/test_responses_api_deployment_check.py diff --git a/tests/litellm/secret_managers/test_get_azure_ad_token_provider.py b/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py similarity index 100% rename from tests/litellm/secret_managers/test_get_azure_ad_token_provider.py rename to tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py diff --git a/tests/litellm/test_constants.py b/tests/test_litellm/test_constants.py similarity index 100% rename from tests/litellm/test_constants.py rename to tests/test_litellm/test_constants.py diff --git a/tests/litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py similarity index 100% rename from tests/litellm/test_cost_calculator.py rename to tests/test_litellm/test_cost_calculator.py diff --git a/tests/litellm/test_logging.py b/tests/test_litellm/test_logging.py similarity index 100% rename from tests/litellm/test_logging.py rename to tests/test_litellm/test_logging.py diff --git a/tests/litellm/test_main.py b/tests/test_litellm/test_main.py similarity index 100% rename from tests/litellm/test_main.py rename to tests/test_litellm/test_main.py diff --git a/tests/litellm/test_router.py b/tests/test_litellm/test_router.py similarity index 99% rename from tests/litellm/test_router.py rename to tests/test_litellm/test_router.py index e9917138f6b..bb266dfd4b2 100644 --- a/tests/litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -298,7 +298,6 @@ async def test_router_amoderation_with_credential_name(mock_amoderation): assert call_kwargs["model"] == "text-moderation-stable" - def test_router_test_team_model(): """ Test that router.test_team_model returns the correct model @@ -319,13 +318,13 @@ def test_router_test_team_model(): result = router.map_team_model(team_model_name="test-model", team_id="test-team") assert result is not None + def test_router_ignore_invalid_deployments(): """ Test that router.ignore_invalid_deployments is set to True """ from litellm.types.router import Deployment - router = litellm.Router( model_list=[ { @@ -348,4 +347,4 @@ def test_router_ignore_invalid_deployments(): ) ) - assert router.get_model_list() == [] \ No newline at end of file + assert router.get_model_list() == [] diff --git a/tests/litellm/test_utils.py b/tests/test_litellm/test_utils.py similarity index 100% rename from tests/litellm/test_utils.py rename to tests/test_litellm/test_utils.py diff --git a/tests/litellm/types/llms/test_types_llms_openai.py b/tests/test_litellm/types/llms/test_types_llms_openai.py similarity index 100% rename from tests/litellm/types/llms/test_types_llms_openai.py rename to tests/test_litellm/types/llms/test_types_llms_openai.py