litellm/tests/test_litellm/conftest.py
yuneng-jiang f6882246d4
test: move tests/test_litellm root and small trees into tests/unit (#43186)
* ci: run the unit_selection.sh shard files on every event instead of only fork pull requests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* ci: rename fork-flag to unit-flag now that it applies on every event

* test: move tests/test_litellm root and small trees into tests/unit

Pure renames, no content changes. Follow-up commits in this PR fix
references, merge the three files that already existed in tests/unit,
keep live-provider tests in tests/test_litellm and wire CI.

* test: carry tests/test_litellm conftest isolation into tests/unit

Callback lists, routing fallbacks, cached HTTP clients, logger state, AWS,
proxy-URL and keychain env, and session-end client cleanup now reset for
unit tests too. The environment isolation owns its MonkeyPatch so a test's
own monkeypatch is undone before the model-cost teardown runs.

* test: merge, split and prune the moved root and small-tree tests

Merge batches/test_batch_utils.py and the chat_completions and messages
dispatch tests into the files that already existed in tests/unit. Keep
the live Gemini interactions tests, the async image-fetch format test and
the OpenAI embedding scorer test in tests/test_litellm since they need
real network or keys. Put test_router.py under tests/unit/test_router so
the existing package no longer shadows it. Delete eight tests the audit
found superseded by stronger ones kept in this move.

* ci: run the moved root and small-tree tests under their legacy flags

Add the misc and responses-caching-types flags to unit_selection.sh and
CircleCI, extend enterprise-routing and mcp-integration, and point the
legacy GHA shards, Makefile, redis-compat workflow, merge smoke manifest
and change classifier at the new paths.

* test: make the new tests/unit directories packages

tests/unit/test_package_layout.py requires every directory to carry an
__init__.py, and without one the moved and retained
test_litellm_responses_bridge.py modules collide on import.

* test: scope the unit socket block to tests/unit in shared sessions

The GHA shards collect the legacy test-path and the unit selection in one
pytest session. The unit conftest's loopback-only block leaked into legacy
modules that reach the network at import. The legacy conftest now lifts the
restriction at collect and setup time, and the unit conftest re-applies it
when collecting its own modules.

* test: give the shard-script tests their own GITHUB_OUTPUT

They only passed where the runner set it. The CircleCI unit job's env
allowlist drops it, so the script's redirect failed there.

* test: point the router and module-deletion checks at tests/unit

router_code_coverage and code_qa_check_tests only searched tests/test_litellm,
so the moved router tests no longer counted. The two silent-experiment tests
the audit deleted were the only direct callers of those methods; they are
replaced with tests that assert the forwarded shadow request and the
recursion guard.

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-09-25 11:30:43 -07:00

664 lines
24 KiB
Python

# conftest.py - IMPROVED VERSION
#
# Key changes:
# 1. Changed module reload from 'module' scope to 'function' scope for better isolation
# 2. Made cache flushing happen per-function instead of per-module
# 3. Removed manual event loop creation (let pytest-asyncio handle it)
# 4. Added proper cleanup in fixtures
# 5. Added worker-specific isolation for parallel execution
import base64
import importlib
import os
from pathlib import Path
from types import SimpleNamespace
import httpx
import pytest
from pytest_socket import _remove_restrictions
import asyncio
import litellm
from litellm import router as litellm_router_module
from litellm import utils as litellm_utils_module
from litellm._logging import ALL_LOGGERS
from litellm.litellm_core_utils.cli_keyring import (
KeyringDiscardsWrites,
KeyringUnreachable,
KeyringUnusable,
SecretErase,
SecretErased,
SecretFound,
SecretMissing,
SecretRead,
SecretStored,
SecretStranded,
SecretWrite,
)
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
def _reset_module_level_aws_auth_caches():
"""
Clear module-level AWS auth state that can survive between tests.
Bedrock/SageMaker handlers are instantiated once at import time and cache
resolved credentials on the handler instance. If a previous test resolves an
invalid or different auth flow, later tests can reuse that cached state and
bypass their local monkeypatched env setup.
"""
for module_name in (
"litellm.main",
"litellm.files.main",
"litellm.rerank_api.main",
"litellm.realtime_api.main",
):
try:
module = importlib.import_module(module_name)
except Exception:
continue
for attr_name in dir(module):
obj = getattr(module, attr_name)
iam_cache = getattr(obj, "iam_cache", None)
if iam_cache is None:
continue
flush_cache = getattr(iam_cache, "flush_cache", None)
if callable(flush_cache):
flush_cache()
try:
import boto3
boto3.DEFAULT_SESSION = None
except Exception:
pass
@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)
monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False)
monkeypatch.delenv("AWS_ROLE_ARN", raising=False)
monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False)
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
monkeypatch.delenv("AWS_REGION_NAME", raising=False)
monkeypatch.delenv("AWS_DEFAULT_REGION", raising=False)
@pytest.fixture(scope="function", autouse=True)
def isolate_host_proxy_base_url(monkeypatch):
"""Prevent a host PROXY_BASE_URL from outranking request-derived URLs during unit tests."""
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
@pytest.fixture(scope="function", autouse=True)
def isolate_host_os_keychain(monkeypatch):
"""Keep any code path that resolves a CLI credential out of the developer's real OS keychain.
Tests that exercise keychain behaviour inject their own vault instead.
"""
monkeypatch.setenv("LITELLM_CLI_DISABLE_KEYRING", "1")
class FakeSecretVault:
"""In-memory stand-in for the OS keychain, injected wherever CLI credential storage is exercised.
`available=False` models a keychain that is locked or has no backend, `writable=False` one that
refuses to store, `erasable=False` one that will not release what it already holds, and `failure`
picks which unusable state those report. `discards=True` is keyring's null backend, which answers
reads and erases like any other yet keeps nothing it is given, so only writes report it.
"""
def __init__(
self,
blob: str | None = None,
*,
available: bool = True,
writable: bool = True,
erasable: bool = True,
discards: bool = False,
failure: KeyringUnusable = KeyringUnreachable(),
) -> None:
self.blob: str | None = blob
self.available: bool = available
self.writable: bool = writable
self.erasable: bool = erasable
self.discards: bool = discards
self.failure: KeyringUnusable = failure
self.reads: int = 0
self.writes: list[str] = []
self.erases: int = 0
def read(self) -> SecretRead:
self.reads += 1
if not self.available:
return self.failure
return SecretMissing() if self.blob is None else SecretFound(self.blob)
def write(self, blob: str) -> SecretWrite:
self.writes.append(blob)
if not (self.available and self.writable):
return self.failure
if self.discards:
return KeyringDiscardsWrites()
self.blob = blob
return SecretStored()
def erase(self) -> SecretErase:
self.erases += 1
if not self.available:
return self.failure
if not self.erasable:
return SecretStranded() if self.blob is not None else SecretErased()
self.blob = None
return SecretErased()
@pytest.fixture
def secret_vault_factory():
"""Build FakeSecretVault instances; see its docstring for the failure modes it can model."""
return FakeSecretVault
@pytest.fixture
def local_model_cost_map(monkeypatch):
"""Force the bundled in-repo cost map so capability and pricing assertions do not
depend on the network-fetched ``main`` copy, which lags this branch until merge.
``get_model_info`` is lru_cached, so swapping ``model_cost`` is not enough on its
own; clear on the way in and out so entries warmed against either map never leak
across tests."""
original_model_cost = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.get_model_info.cache_clear()
try:
yield
finally:
litellm.model_cost = original_model_cost
litellm.get_model_info.cache_clear()
@pytest.fixture
def local_beta_headers_config(monkeypatch):
"""Pin the bundled ``anthropic_beta_headers_config.json`` so beta header assertions
do not depend on the network-fetched copy or on what earlier tests left cached."""
from litellm.anthropic_beta_headers_manager import reload_beta_headers_config
monkeypatch.setenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", "True")
reload_beta_headers_config()
try:
yield
finally:
monkeypatch.delenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", raising=False)
reload_beta_headers_config()
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)
def isolate_litellm_state():
"""
Per-function isolation fixture (changed from module scope).
This ensures better isolation when running tests in parallel:
- Each test function gets a clean litellm state
- Cache is flushed before each test
- No module reloading during parallel execution
Note: Module reloading at function scope is safer for parallel execution
but adds overhead. Consider removing reload entirely if tests can work without it.
"""
# Get worker ID if running with pytest-xdist
worker_id = os.environ.get("PYTEST_XDIST_WORKER", "master")
# Store original callback state (all callback lists)
original_state = {}
if hasattr(litellm, "callbacks"):
original_state["callbacks"] = (
litellm.callbacks.copy() if litellm.callbacks else []
)
if hasattr(litellm, "success_callback"):
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
if hasattr(litellm, "model_fallbacks"):
original_state["model_fallbacks"] = litellm.model_fallbacks
# Store transport/network globals — many tests set these without restoring,
# causing subsequent tests to get None from _create_async_transport()
for _attr in ("disable_aiohttp_transport", "force_ipv4"):
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)
original_runtime_registered_model_cost = {
model_key: dict(model_value)
for model_key, model_value in litellm_utils_module._runtime_registered_model_cost.items()
}
original_live_routers = set(litellm_router_module._live_routers)
# 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()
_reset_module_level_aws_auth_caches()
# litellm.get_model_info() memoizes ModelInfo built from litellm.model_cost, so a
# test that rebinds the cost map leaves later tests pricing against the old map.
litellm_utils_module._invalidate_model_cost_lowercase_map()
# Clear all callback lists to prevent cross-test contamination
if hasattr(litellm, "callbacks"):
litellm.callbacks = []
if hasattr(litellm, "success_callback"):
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()
_reset_module_level_aws_auth_caches()
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)
litellm_utils_module._runtime_registered_model_cost.clear()
litellm_utils_module._runtime_registered_model_cost.update(original_runtime_registered_model_cost)
litellm_utils_module._invalidate_model_cost_lowercase_map()
for _router in tuple(litellm_router_module._live_routers):
litellm_router_module._live_routers.discard(_router)
for _router in original_live_routers:
litellm_router_module._live_routers.add(_router)
# 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():
"""
Module-scoped setup/teardown for heavy initialization.
Use this sparingly - most state should be handled by isolate_litellm_state.
Only reload modules here if absolutely necessary.
"""
import litellm
# Only reload if NOT running in parallel (module reload + parallel = bad)
worker_id = os.environ.get("PYTEST_XDIST_WORKER", None)
if worker_id is None:
# Single process mode - safe to reload
importlib.reload(litellm)
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
import litellm.proxy.proxy_server
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
# Flush cache after reload (prevents stale client instances)
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
print(f"[conftest] Module setup complete (worker: {worker_id or 'master'})")
yield
# Teardown - no need to manually manage event loops with pytest-asyncio auto mode
print(f"[conftest] Module teardown complete (worker: {worker_id or 'master'})")
def pytest_collectstart():
_remove_restrictions()
def pytest_runtest_setup():
_remove_restrictions()
def pytest_collection_modifyitems(config, items):
"""
Customize test collection order.
- Separate tests marked with 'no_parallel' from parallelizable tests
- Sort custom_logger tests first (they tend to interfere with other tests)
"""
# Separate no_parallel tests
no_parallel_tests = [
item
for item in items
if any(mark.name == "no_parallel" for mark in item.iter_markers())
]
# Separate custom_logger tests
custom_logger_tests = [
item
for item in items
if "custom_logger" in item.parent.name and item not in no_parallel_tests
]
# Everything else
other_tests = [
item
for item in items
if item not in no_parallel_tests and item not in custom_logger_tests
]
# Sort each group
custom_logger_tests.sort(key=lambda x: x.name)
other_tests.sort(key=lambda x: x.name)
no_parallel_tests.sort(key=lambda x: x.name)
# Reorder: custom_logger first (isolated), then other tests, then no_parallel tests last
items[:] = custom_logger_tests + other_tests + no_parallel_tests
def pytest_configure(config):
"""
Configure pytest with custom settings.
"""
# Add marker for flaky tests (for documentation purposes)
config.addinivalue_line(
"markers", "flaky: mark test as potentially flaky (should use --reruns)"
)
# Detect if running in CI
is_ci = os.environ.get("CI") == "true" or os.environ.get("LITELLM_CI") == "true"
if is_ci:
print("[conftest] Running in CI mode - enabling stricter test isolation")
# Optional: Add a fixture for tests that need even stricter isolation
@pytest.fixture
def strict_isolation():
"""
Use this fixture for tests that need extra strict isolation.
Example:
def test_something(strict_isolation):
# Test code with guaranteed clean state
pass
"""
# Force flush all caches
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
# Reset all global state
if hasattr(litellm, "disable_aiohttp_transport"):
original_aiohttp = litellm.disable_aiohttp_transport
litellm.disable_aiohttp_transport = False
else:
original_aiohttp = None
if hasattr(litellm, "set_verbose"):
original_verbose = litellm.set_verbose
litellm.set_verbose = False
else:
original_verbose = None
yield
# Restore original state
if original_aiohttp is not None:
litellm.disable_aiohttp_transport = original_aiohttp
if original_verbose is not None:
litellm.set_verbose = original_verbose
# 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())
ONE_PIXEL_PNG = base64.b64decode(
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
)
@pytest.fixture
def async_only_image_fetch(monkeypatch):
from litellm.litellm_core_utils.prompt_templates import factory, image_handling
from litellm.llms.gemini.chat import transformation as gemini_chat_transformation
fetch = SimpleNamespace(
fetched=[],
base64_png=base64.b64encode(ONE_PIXEL_PNG).decode(),
data_url="data:image/png;base64," + base64.b64encode(ONE_PIXEL_PNG).decode(),
)
def forbid_sync_fetch(client, url, **kwargs):
raise litellm.ImageFetchError(f"sync image fetch ran on the event loop: {url}")
async def serve_png(client, url, **kwargs):
fetch.fetched.append(url)
return httpx.Response(
200,
content=ONE_PIXEL_PNG,
headers={"content-type": "image/png"},
request=httpx.Request("GET", url),
)
def forbid_sync_convert(url, *args, **kwargs):
if url.startswith(("http://", "https://")):
raise litellm.ImageFetchError(f"sync convert_url_to_base64 ran on the request path: {url}")
return url
monkeypatch.setattr(image_handling, "safe_get", forbid_sync_fetch)
monkeypatch.setattr(image_handling, "async_safe_get", serve_png)
for module in (image_handling, factory, gemini_chat_transformation):
monkeypatch.setattr(module, "convert_url_to_base64", forbid_sync_convert)
return fetch