diff --git a/tests/test_litellm/conftest.py b/tests/test_litellm/conftest.py index 4421d227f4e..7b3c0b27842 100644 --- a/tests/test_litellm/conftest.py +++ b/tests/test_litellm/conftest.py @@ -10,6 +10,7 @@ import importlib import os import sys +from pathlib import Path import pytest sys.path.insert( @@ -18,6 +19,75 @@ sys.path.insert( import asyncio import litellm +from litellm._logging import ALL_LOGGERS +from litellm.litellm_core_utils.prompt_templates import ( + image_handling as image_handling_module, +) +from litellm.llms.custom_httpx.async_client_cleanup import ( + close_litellm_async_clients, +) +from litellm.proxy.db import tool_registry_writer as tool_registry_writer_module + + +@pytest.fixture(scope="session") +def isolated_aws_credentials_dir(tmp_path_factory): + aws_dir = tmp_path_factory.mktemp("aws-config") + credentials_file = Path(aws_dir) / "credentials" + config_file = Path(aws_dir) / "config" + credentials_file.write_text("", encoding="utf-8") + config_file.write_text("", encoding="utf-8") + return { + "credentials": str(credentials_file), + "config": str(config_file), + } + + +@pytest.fixture(scope="function", autouse=True) +def isolate_host_aws_config(monkeypatch, isolated_aws_credentials_dir): + """Prevent botocore from reading host AWS profiles during unit tests.""" + monkeypatch.setenv( + "AWS_SHARED_CREDENTIALS_FILE", isolated_aws_credentials_dir["credentials"] + ) + monkeypatch.setenv("AWS_CONFIG_FILE", isolated_aws_credentials_dir["config"]) + monkeypatch.setenv("AWS_EC2_METADATA_DISABLED", "true") + monkeypatch.delenv("AWS_PROFILE", raising=False) + monkeypatch.delenv("AWS_DEFAULT_PROFILE", raising=False) + monkeypatch.delenv("AWS_CONTAINER_CREDENTIALS_FULL_URI", raising=False) + monkeypatch.delenv("AWS_CONTAINER_CREDENTIALS_RELATIVE_URI", raising=False) + + +def _run_coroutine_if_needed(result): + if not asyncio.iscoroutine(result): + return + + try: + asyncio.run(result) + except RuntimeError: + # If pytest-asyncio already has a running loop, best-effort scheduling is + # still better than leaking the client entirely. + try: + loop = asyncio.get_running_loop() + except RuntimeError: + return + if loop.is_running(): + loop.create_task(result) + except Exception: + pass + + +def _close_handler_if_needed(handler): + if handler is None: + return + + close_fn = getattr(handler, "close", None) + if not callable(close_fn): + return + + try: + result = close_fn() + _run_coroutine_if_needed(result) + except Exception: + pass @pytest.fixture(scope="function", autouse=True) @@ -44,10 +114,14 @@ def isolate_litellm_state(): original_state['success_callback'] = litellm.success_callback.copy() if litellm.success_callback else [] if hasattr(litellm, 'failure_callback'): original_state['failure_callback'] = litellm.failure_callback.copy() if litellm.failure_callback else [] + if hasattr(litellm, 'input_callback'): + original_state['input_callback'] = litellm.input_callback.copy() if litellm.input_callback else [] if hasattr(litellm, '_async_success_callback'): original_state['_async_success_callback'] = litellm._async_success_callback.copy() if litellm._async_success_callback else [] if hasattr(litellm, '_async_failure_callback'): original_state['_async_failure_callback'] = litellm._async_failure_callback.copy() if litellm._async_failure_callback else [] + if hasattr(litellm, '_async_input_callback'): + original_state['_async_input_callback'] = litellm._async_input_callback.copy() if litellm._async_input_callback else [] # Store routing globals — leaked model_fallbacks causes tests to route # through async_completion_with_fallbacks / Router, bypassing HTTP mocks @@ -60,9 +134,68 @@ def isolate_litellm_state(): if hasattr(litellm, _attr): original_state[_attr] = getattr(litellm, _attr) + # Store request-mapping globals that are frequently mutated in tests. + if hasattr(litellm, "drop_params"): + original_state["drop_params"] = litellm.drop_params + if hasattr(litellm, "cache"): + original_state["cache"] = litellm.cache + + # Store secret-manager globals. Several tests swap these out, which changes + # get_secret() behavior for later env-driven tests (for example Redis config). + for _attr in ("secret_manager_client", "_key_management_system", "_key_management_settings"): + if hasattr(litellm, _attr): + original_state[_attr] = getattr(litellm, _attr) + + # Store other commonly-mutated LiteLLM globals that affect provider routing, + # auth, and request shaping during larger suite runs. + for _attr in ( + "api_base", + "num_retries", + "modify_params", + "ssl_verify", + "credential_list", + "model_group_settings", + "default_internal_user_params", + "default_team_params", + "prometheus_emit_stream_label", + "vector_store_registry", + "model_cost", + "cost_margin_config", + "cost_discount_config", + "disable_hf_tokenizer_download", + "disable_copilot_system_to_assistant", + "cohere_models", + "anthropic_models", + "token_counter", + "initialized_langfuse_clients", + ): + if hasattr(litellm, _attr): + original_state[_attr] = getattr(litellm, _attr) + + # Store LiteLLM logger state. Some tests reconfigure handlers/propagation for + # JSON logging and do not restore them, which breaks later caplog-based tests. + logger_state = {} + for logger in ALL_LOGGERS: + logger_state[logger.name] = { + "level": logger.level, + "disabled": logger.disabled, + "propagate": logger.propagate, + "handlers": list(logger.handlers), + "filters": list(logger.filters), + } + + # Store singleton registries that are lazily initialized during tests and + # can change endpoint behavior later in the suite. + original_tool_policy_registry = tool_registry_writer_module._tool_policy_registry + had_module_level_client = "module_level_client" in litellm.__dict__ + had_module_level_aclient = "module_level_aclient" in litellm.__dict__ + original_module_level_client = litellm.__dict__.get("module_level_client") + original_module_level_aclient = litellm.__dict__.get("module_level_aclient") + # Flush cache before test (critical for respx mocks) if hasattr(litellm, "in_memory_llm_clients_cache"): litellm.in_memory_llm_clients_cache.flush_cache() + image_handling_module.in_memory_cache.flush_cache() # Clear all callback lists to prevent cross-test contamination if hasattr(litellm, 'callbacks'): @@ -71,26 +204,63 @@ def isolate_litellm_state(): litellm.success_callback = [] if hasattr(litellm, 'failure_callback'): litellm.failure_callback = [] + if hasattr(litellm, 'input_callback'): + litellm.input_callback = [] if hasattr(litellm, '_async_success_callback'): litellm._async_success_callback = [] if hasattr(litellm, '_async_failure_callback'): litellm._async_failure_callback = [] + if hasattr(litellm, '_async_input_callback'): + litellm._async_input_callback = [] # Clear routing globals if hasattr(litellm, 'model_fallbacks'): litellm.model_fallbacks = None + if hasattr(litellm, "cache"): + litellm.cache = None + litellm.__dict__.pop("module_level_client", None) + litellm.__dict__.pop("module_level_aclient", None) + tool_registry_writer_module._tool_policy_registry = None yield # Cleanup after test if hasattr(litellm, "in_memory_llm_clients_cache"): litellm.in_memory_llm_clients_cache.flush_cache() + image_handling_module.in_memory_cache.flush_cache() + current_module_level_client = litellm.__dict__.get("module_level_client") + current_module_level_aclient = litellm.__dict__.get("module_level_aclient") # Restore all callback lists to original state for attr_name, original_value in original_state.items(): if hasattr(litellm, attr_name): setattr(litellm, attr_name, original_value) + # Restore logger configuration mutated by logging-focused tests. + for logger in ALL_LOGGERS: + original_logger_state = logger_state.get(logger.name) + if original_logger_state is None: + continue + logger.setLevel(original_logger_state["level"]) + logger.disabled = original_logger_state["disabled"] + logger.propagate = original_logger_state["propagate"] + logger.handlers = list(original_logger_state["handlers"]) + logger.filters = list(original_logger_state["filters"]) + + tool_registry_writer_module._tool_policy_registry = original_tool_policy_registry + if current_module_level_client is not original_module_level_client: + _close_handler_if_needed(current_module_level_client) + if current_module_level_aclient is not original_module_level_aclient: + _close_handler_if_needed(current_module_level_aclient) + if had_module_level_client: + litellm.__dict__["module_level_client"] = original_module_level_client + else: + litellm.__dict__.pop("module_level_client", None) + if had_module_level_aclient: + litellm.__dict__["module_level_aclient"] = original_module_level_aclient + else: + litellm.__dict__.pop("module_level_aclient", None) + @pytest.fixture(scope="module", autouse=True) def setup_and_teardown(): @@ -220,3 +390,16 @@ def strict_isolation(): # Final cache flush if hasattr(litellm, "in_memory_llm_clients_cache"): litellm.in_memory_llm_clients_cache.flush_cache() + + +def pytest_sessionfinish(session, exitstatus): + """Close any globally cached HTTP clients so xdist workers exit cleanly.""" + _close_handler_if_needed(litellm.__dict__.get("module_level_client")) + _close_handler_if_needed(litellm.__dict__.get("module_level_aclient")) + litellm.__dict__.pop("module_level_client", None) + litellm.__dict__.pop("module_level_aclient", None) + _close_handler_if_needed(getattr(litellm, "base_llm_aiohttp_handler", None)) + _close_handler_if_needed(getattr(litellm, "httpx_client", None)) + _close_handler_if_needed(getattr(litellm, "aclient", None)) + _close_handler_if_needed(getattr(litellm, "client", None)) + _run_coroutine_if_needed(close_litellm_async_clients()) diff --git a/tests/test_litellm/integrations/gitlab/test_gitlab_prompt_manager.py b/tests/test_litellm/integrations/gitlab/test_gitlab_prompt_manager.py index d623dba0c34..a11c2cd4fa8 100644 --- a/tests/test_litellm/integrations/gitlab/test_gitlab_prompt_manager.py +++ b/tests/test_litellm/integrations/gitlab/test_gitlab_prompt_manager.py @@ -87,7 +87,7 @@ def test_gitlab_client_missing_required_fields(): # GitLabClient: get_file_content # ----------------------- -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.get") +@patch("litellm.integrations.gitlab.gitlab_client.HTTPHandler.get") def test_gitlab_client_get_file_content_raw_success(mock_get): """Successful file content retrieval via RAW endpoint.""" mock_response = MagicMock() @@ -104,7 +104,7 @@ def test_gitlab_client_get_file_content_raw_success(mock_get): mock_get.assert_called_once() -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.get") +@patch("litellm.integrations.gitlab.gitlab_client.HTTPHandler.get") def test_gitlab_client_get_file_content_raw_404_fallback_json_base64(mock_get): """When RAW returns 404, fallback to JSON endpoint and decode base64 content.""" import base64 @@ -136,7 +136,7 @@ def test_gitlab_client_get_file_content_raw_404_fallback_json_base64(mock_get): assert content == "json-content" -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.get") +@patch("litellm.integrations.gitlab.gitlab_client.HTTPHandler.get") def test_gitlab_client_get_file_content_not_found(mock_get): """File not found returns None.""" # Simulate RAW 404 and JSON 404 @@ -152,7 +152,7 @@ def test_gitlab_client_get_file_content_not_found(mock_get): assert content is None -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.get") +@patch("litellm.integrations.gitlab.gitlab_client.HTTPHandler.get") def test_gitlab_client_get_file_content_access_denied(mock_get): """403 raises a helpful message.""" import httpx @@ -168,7 +168,7 @@ def test_gitlab_client_get_file_content_access_denied(mock_get): client.get_file_content("test.prompt") -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.get") +@patch("litellm.integrations.gitlab.gitlab_client.HTTPHandler.get") def test_gitlab_client_get_file_content_auth_failed(mock_get): """401 raises auth error.""" import httpx @@ -186,7 +186,7 @@ def test_gitlab_client_get_file_content_auth_failed(mock_get): # GitLabClient: list_files # ----------------------- -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.get") +@patch("litellm.integrations.gitlab.gitlab_client.HTTPHandler.get") def test_gitlab_client_list_files_success(mock_get): """List .prompt files via repository tree API.""" mock_response = MagicMock() @@ -817,4 +817,3 @@ def test_cache_get_by_file_returns_exact_entry(mock_pm_cls, fake_managers): assert alpha and alpha["id"] == "alpha" assert beta and beta["id"] == "nested/beta" - diff --git a/tests/test_litellm/proxy/client/test_chat.py b/tests/test_litellm/proxy/client/test_chat.py index 4de58723761..b8e55c45502 100644 --- a/tests/test_litellm/proxy/client/test_chat.py +++ b/tests/test_litellm/proxy/client/test_chat.py @@ -1,20 +1,53 @@ -import os +import importlib +import importlib.util +from importlib.machinery import PathFinder +import site import sys import pytest import requests - -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system path - - -import responses - from litellm.proxy.client.chat import ChatClient from litellm.proxy.client.exceptions import UnauthorizedError +def _load_http_mocking_responses(): + """Load the third-party `responses` package even if test collection creates + a top-level `responses` namespace package from `tests/test_litellm/responses`. + """ + module = importlib.import_module("responses") + if hasattr(module, "activate"): + return module + + for module_name in list(sys.modules): + if module_name == "responses" or module_name.startswith("responses."): + sys.modules.pop(module_name, None) + + search_paths = [] + try: + search_paths.extend(site.getsitepackages()) + except AttributeError: + pass + user_site = site.getusersitepackages() + if isinstance(user_site, str): + search_paths.append(user_site) + else: + search_paths.extend(user_site) + + spec = PathFinder.find_spec("responses", search_paths) + if spec is None or spec.loader is None: + raise ImportError("Unable to load the third-party `responses` package") + module = importlib.util.module_from_spec(spec) + sys.modules["responses"] = module + spec.loader.exec_module(module) + + if not hasattr(module, "activate"): + raise ImportError("Unable to load the third-party `responses` package") + return module + + +responses = _load_http_mocking_responses() + + @pytest.fixture def base_url(): return "http://localhost:8000" diff --git a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py index 25ff6f89427..8aeb1009101 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py @@ -11,7 +11,7 @@ from litellm.proxy._types import ( LitellmUserRoles, ProxyException, ) -from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.management_endpoints.customer_endpoints import router app = FastAPI() @@ -42,11 +42,14 @@ def mock_prisma_client(): @pytest.fixture def mock_user_api_key_auth(): - with patch("litellm.proxy.proxy_server.user_api_key_auth") as mock: - mock.return_value = UserAPIKeyAuth( - user_id="test-user", user_role=LitellmUserRoles.PROXY_ADMIN - ) - yield mock + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user", user_role=LitellmUserRoles.PROXY_ADMIN + ) + try: + yield + finally: + app.dependency_overrides = original_overrides def test_update_customer_success(mock_prisma_client, mock_user_api_key_auth): diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 349fe76ed71..6c6ea11bf90 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -348,7 +348,10 @@ class TestProxyInitializationHelpers: }, ), patch( "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" - ) as mock_get_args: + ) as mock_get_args, patch( + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._is_port_in_use", + return_value=False, + ): mock_get_args.return_value = { "app": "litellm.proxy.proxy_server:app", "host": "localhost", diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index ce5873f5063..c241813d08a 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -39,6 +39,12 @@ def add_api_keys_to_env(monkeypatch): monkeypatch.setenv("AWS_ACCESS_KEY_ID", "my-fake-aws-access-key-id") monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "my-fake-aws-secret-access-key") monkeypatch.setenv("AWS_REGION", "us-east-1") + # Keep these transformation tests on the simple access-key path. A leaked + # session token or role/web-identity env var pushes Bedrock auth down a + # different branch and fails before the mocked HTTP client is exercised. + monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False) + monkeypatch.delenv("AWS_ROLE_ARN", raising=False) + monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False) @pytest.fixture