From 5db6aef8344dee6a43d2fc3fd3716d82d6471523 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sun, 15 Mar 2026 23:33:21 -0700 Subject: [PATCH] [Fix] Restore xdist test isolation: capture true defaults and poll cooldowns The revert of 9711e3adfe left xdist tests without proper state isolation. Module-level assignments like `litellm.num_retries = 3` in 12+ test files pollute shared globals, and the fixture was saving/restoring contaminated values instead of resetting to true defaults. - Capture true litellm defaults at conftest import time and reset before each test (local_testing + llm_translation) - Make llm_translation/conftest.py xdist-safe (skip reload under xdist, add state isolation) - Replace asyncio.sleep(2) with polling in cooldown handler tests Co-Authored-By: Claude Opus 4.6 --- tests/llm_translation/conftest.py | 77 +++++++++++++++++-- tests/local_testing/conftest.py | 71 +++++++++++------ .../test_router_cooldown_handlers.py | 12 ++- 3 files changed, 126 insertions(+), 34 deletions(-) diff --git a/tests/llm_translation/conftest.py b/tests/llm_translation/conftest.py index 97edb4c023c..46a3a31771f 100644 --- a/tests/llm_translation/conftest.py +++ b/tests/llm_translation/conftest.py @@ -1,9 +1,14 @@ # conftest.py +# +# xdist-compatible test isolation for llm_translation tests. +# Mirrors the pattern in tests/local_testing/conftest.py: +# - Function-scoped fixture resets litellm globals to true defaults +# - Module-scoped reload only in single-process mode import importlib import os import sys -import asyncio + import pytest sys.path.insert( @@ -13,6 +18,21 @@ import litellm import asyncio +# --------------------------------------------------------------------------- +# Capture TRUE defaults at conftest import time (before test modules pollute). +# --------------------------------------------------------------------------- +_SCALAR_DEFAULTS = { + "num_retries": getattr(litellm, "num_retries", None), + "set_verbose": getattr(litellm, "set_verbose", False), + "cache": getattr(litellm, "cache", None), + "allowed_fails": getattr(litellm, "allowed_fails", 3), + "disable_aiohttp_transport": getattr(litellm, "disable_aiohttp_transport", False), + "force_ipv4": getattr(litellm, "force_ipv4", False), + "drop_params": getattr(litellm, "drop_params", None), + "modify_params": getattr(litellm, "modify_params", False), +} + + @pytest.fixture(scope="session") def event_loop(): try: @@ -29,20 +49,63 @@ def setup_and_teardown(event_loop): # Add event_loop as a dependency sys.path.insert(0, os.path.abspath("../..")) import litellm - from litellm import Router - from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER - # flush all logs - asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue()) + # ---- Save current state (for teardown restore) ---- + original_state = {} + for attr in ( + "callbacks", + "success_callback", + "failure_callback", + "_async_success_callback", + "_async_failure_callback", + ): + if hasattr(litellm, attr): + val = getattr(litellm, attr) + original_state[attr] = val.copy() if val else [] - importlib.reload(litellm) + for attr in _SCALAR_DEFAULTS: + if hasattr(litellm, attr): + original_state[attr] = getattr(litellm, attr) + + # ---- Reset to true defaults before the test ---- + worker_id = os.environ.get("PYTEST_XDIST_WORKER", None) + if worker_id is None: + # Single-process mode: reload for full reset + from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue()) + importlib.reload(litellm) + else: + # xdist mode: reset globals without reload + for attr in ( + "callbacks", + "success_callback", + "failure_callback", + "_async_success_callback", + "_async_failure_callback", + ): + if hasattr(litellm, attr): + setattr(litellm, attr, []) + + for attr, default_val in _SCALAR_DEFAULTS.items(): + if hasattr(litellm, attr): + setattr(litellm, attr, default_val) + + if hasattr(litellm, "in_memory_llm_clients_cache"): + litellm.in_memory_llm_clients_cache.flush_cache() # Set the event loop from the fixture asyncio.set_event_loop(event_loop) - print(litellm) yield + # ---- Teardown ---- + if hasattr(litellm, "in_memory_llm_clients_cache"): + litellm.in_memory_llm_clients_cache.flush_cache() + + for attr, original_value in original_state.items(): + if hasattr(litellm, attr): + setattr(litellm, attr, original_value) + # Clean up any pending tasks pending = asyncio.all_tasks(event_loop) for task in pending: diff --git a/tests/local_testing/conftest.py b/tests/local_testing/conftest.py index b858154e2fb..a6b31caff22 100644 --- a/tests/local_testing/conftest.py +++ b/tests/local_testing/conftest.py @@ -4,6 +4,12 @@ # Pattern matches tests/test_litellm/conftest.py: # - Function-scoped fixture saves/restores litellm globals (no reload) # - Module-scoped fixture reloads only in single-process mode +# +# IMPORTANT: True defaults are captured at conftest import time (before any +# test module can pollute them via module-level assignments like +# `litellm.num_retries = 3`). The function-scoped fixture resets globals to +# these true defaults before every test, preventing cross-test contamination +# under xdist where module reload is skipped. import importlib import os @@ -16,16 +22,40 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm +# --------------------------------------------------------------------------- +# Capture TRUE defaults at conftest import time. This runs before any test +# module's top-level code (e.g. `litellm.num_retries = 3`) executes, so +# the values here are guaranteed to be the real package defaults. +# --------------------------------------------------------------------------- +_SCALAR_DEFAULTS = { + "num_retries": getattr(litellm, "num_retries", None), + "num_retries_per_request": getattr(litellm, "num_retries_per_request", None), + "request_timeout": getattr(litellm, "request_timeout", None), + "set_verbose": getattr(litellm, "set_verbose", False), + "cache": getattr(litellm, "cache", None), + "allowed_fails": getattr(litellm, "allowed_fails", 3), + "default_fallbacks": getattr(litellm, "default_fallbacks", None), + "enable_azure_ad_token_refresh": getattr(litellm, "enable_azure_ad_token_refresh", None), + "tag_budget_config": getattr(litellm, "tag_budget_config", None), + "model_cost": getattr(litellm, "model_cost", None), + "token_counter": getattr(litellm, "token_counter", None), + "disable_aiohttp_transport": getattr(litellm, "disable_aiohttp_transport", False), + "force_ipv4": getattr(litellm, "force_ipv4", False), + "drop_params": getattr(litellm, "drop_params", None), + "modify_params": getattr(litellm, "modify_params", False), +} + @pytest.fixture(scope="function", autouse=True) def isolate_litellm_state(): """ Per-function isolation fixture. - Saves and restores litellm callback/global state so tests don't leak - side effects. Works safely under pytest-xdist parallel execution. + Resets litellm globals to their true defaults before each test and + restores them afterward, so tests don't leak side effects. + Works safely under pytest-xdist parallel execution. """ - # Save original callback state + # ---- Save current callback state (for teardown restore) ---- original_state = {} for attr in ( "callbacks", @@ -38,38 +68,23 @@ def isolate_litellm_state(): val = getattr(litellm, attr) original_state[attr] = val.copy() if val else [] - # Save other globals that tests commonly mutate - for attr in ( - "set_verbose", - "cache", - "num_retries", - "num_retries_per_request", - "request_timeout", - "default_fallbacks", - "enable_azure_ad_token_refresh", - "tag_budget_config", - "model_cost", - "token_counter", - ): - if hasattr(litellm, attr): - original_state[attr] = getattr(litellm, attr) - - # Save rules that tests may set (e.g. test_rules.py) + # Save list-type globals for attr in ("pre_call_rules", "post_call_rules"): if hasattr(litellm, attr): val = getattr(litellm, attr) original_state[attr] = val.copy() if val else [] - # Save transport/network globals - for attr in ("disable_aiohttp_transport", "force_ipv4"): + # Save scalar globals + for attr in _SCALAR_DEFAULTS: if hasattr(litellm, attr): original_state[attr] = getattr(litellm, attr) - # Flush cache before test + # ---- Reset to true defaults before the test ---- + # Flush HTTP client cache if hasattr(litellm, "in_memory_llm_clients_cache"): litellm.in_memory_llm_clients_cache.flush_cache() - # Clear callbacks and rules before test + # Clear callbacks and rules for attr in ( "callbacks", "success_callback", @@ -82,9 +97,15 @@ def isolate_litellm_state(): if hasattr(litellm, attr): setattr(litellm, attr, []) + # Reset scalar globals to true defaults (prevents contamination from + # module-level code like `litellm.num_retries = 3` in test files) + for attr, default_val in _SCALAR_DEFAULTS.items(): + if hasattr(litellm, attr): + setattr(litellm, attr, default_val) + yield - # Restore all saved state + # ---- Teardown: restore saved state ---- if hasattr(litellm, "in_memory_llm_clients_cache"): litellm.in_memory_llm_clients_cache.flush_cache() diff --git a/tests/local_testing/test_router_cooldown_handlers.py b/tests/local_testing/test_router_cooldown_handlers.py index 012dcb5808e..7be8289abf1 100644 --- a/tests/local_testing/test_router_cooldown_handlers.py +++ b/tests/local_testing/test_router_cooldown_handlers.py @@ -376,7 +376,11 @@ async def test_single_deployment_cooldown_with_allowed_fails(): except litellm.Timeout: pass - await asyncio.sleep(2) + # Poll until the mock is called (or timeout) + for _ in range(40): + if mock_client.call_count >= 1: + break + await asyncio.sleep(0.1) mock_client.assert_called_once() @@ -426,7 +430,11 @@ async def test_single_deployment_cooldown_with_allowed_fail_policy(): except litellm.Timeout: pass - await asyncio.sleep(2) + # Poll until the mock is called (or timeout) + for _ in range(40): + if mock_client.call_count >= 1: + break + await asyncio.sleep(0.1) mock_client.assert_called_once()