test: harden test_litellm isolation

This commit is contained in:
user 2026-04-02 08:23:33 +00:00
parent 2ad24218d0
commit a7c11356f4
No known key found for this signature in database
6 changed files with 250 additions and 23 deletions

View file

@ -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())

View file

@ -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"

View file

@ -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"

View file

@ -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):

View file

@ -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",

View file

@ -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