mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
test: harden test_litellm isolation
This commit is contained in:
parent
2ad24218d0
commit
a7c11356f4
6 changed files with 250 additions and 23 deletions
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue