Merge pull request #20595 from jquinter/fix/test-parallelization-isolation

fix: improve test isolation for parallel execution
This commit is contained in:
jquinter 2026-02-17 21:46:31 -03:00 • committed by GitHub
commit cca4a8699a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 67 additions and 44 deletions

View file

@ -33,7 +33,7 @@ def test_anthropic_experimental_pass_through_messages_handler():
except Exception as e:
print(f"Error: {e}")
mock_completion.assert_called_once()
mock_completion.call_args.kwargs["api_key"] == "test-api-key"
assert mock_completion.call_args.kwargs["api_key"] == "test-api-key"
def test_anthropic_experimental_pass_through_messages_handler_dynamic_api_key_and_api_base_and_custom_values():
@ -57,9 +57,9 @@ def test_anthropic_experimental_pass_through_messages_handler_dynamic_api_key_an
except Exception as e:
print(f"Error: {e}")
mock_completion.assert_called_once()
mock_completion.call_args.kwargs["api_key"] == "test-api-key"
mock_completion.call_args.kwargs["api_base"] == "test-api-base"
mock_completion.call_args.kwargs["custom_key"] == "custom_value"
assert mock_completion.call_args.kwargs["api_key"] == "test-api-key"
assert mock_completion.call_args.kwargs["api_base"] == "test-api-base"
assert mock_completion.call_args.kwargs["custom_key"] == "custom_value"
def test_anthropic_experimental_pass_through_messages_handler_custom_llm_provider():

View file

@ -1,3 +1,4 @@
import importlib
import json
import os
import sys
@ -15,7 +16,22 @@ MOCK_EMBEDDING_RESPONSE = [[0.1, 0.2, 0.3, 0.4, 0.5]]
@pytest.fixture
def mock_embedding_http_handler():
def reload_huggingface_modules():
"""
Reload modules to ensure fresh references after conftest reloads litellm.
This ensures the HTTPHandler class being patched is the same one used by
the embedding handler during parallel test execution.
"""
import litellm.llms.custom_httpx.http_handler as http_handler_module
import litellm.llms.huggingface.embedding.handler as hf_embedding_handler_module
importlib.reload(http_handler_module)
importlib.reload(hf_embedding_handler_module)
yield
@pytest.fixture
def mock_embedding_http_handler(reload_huggingface_modules):
"""Fixture to mock the HTTP handler for embedding tests"""
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
mock_response = MagicMock()
@ -27,7 +43,7 @@ def mock_embedding_http_handler():
@pytest.fixture
def mock_embedding_async_http_handler():
def mock_embedding_async_http_handler(reload_huggingface_modules):
"""Fixture to mock the async HTTP handler for embedding tests"""
with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", new_callable=AsyncMock) as mock_post:
mock_response = MagicMock()

View file

@ -2,6 +2,7 @@
Integration tests for Vertex AI rerank functionality.
These tests demonstrate end-to-end usage of the Vertex AI rerank feature.
"""
import importlib
import os
from unittest.mock import MagicMock, patch
@ -13,7 +14,14 @@ from litellm.llms.vertex_ai.rerank.transformation import VertexAIRerankConfig
class TestVertexAIRerankIntegration:
def setup_method(self):
self.config = VertexAIRerankConfig()
# Reload modules to ensure fresh references after conftest reloads litellm.
# This ensures the class being patched is the same one used by the tests.
import litellm.llms.vertex_ai.rerank.transformation as rerank_transformation_module
importlib.reload(rerank_transformation_module)
# Re-import after reload to get the fresh class
from litellm.llms.vertex_ai.rerank.transformation import VertexAIRerankConfig as FreshConfig
self.config = FreshConfig()
self.model = "semantic-ranker-default@latest"
@patch('litellm.llms.vertex_ai.rerank.transformation.VertexAIRerankConfig._ensure_access_token')

View file

@ -6,6 +6,7 @@ and following LiteLLM testing patterns and best practices.
"""
# Standard library imports
import importlib
import os
import sys
from typing import Dict
@ -49,16 +50,14 @@ def setup_and_teardown():
to speed up testing by removing callbacks being chained.
"""
import asyncio
import importlib
import sys
global litellm
# Reload litellm to ensure clean state
# During parallel test execution, another worker might have removed litellm from sys.modules
# so we need to ensure it's imported before reloading
if "litellm" not in sys.modules:
import litellm as _litellm
else:
importlib.reload(litellm)
# Always import then reload to ensure fresh state
# This handles both cases uniformly:
# 1. litellm not in sys.modules (parallel worker removed it)
# 2. litellm already imported (normal case)
_module = importlib.import_module("litellm")
litellm = importlib.reload(_module)
# Set up async loop
loop = asyncio.get_event_loop_policy().new_event_loop()

View file

@ -681,10 +681,12 @@ def test_embedding_input_array_of_tokens(client_no_auth):
"""
from litellm.proxy import proxy_server
# Apply the mock AFTER client_no_auth fixture has initialized the router
# This avoids issues with llm_router being None during parallel test execution
if proxy_server.llm_router is None:
pytest.skip("llm_router not initialized - skipping test")
# The client_no_auth fixture should initialize the router
# Assert this to catch any router initialization regressions
assert proxy_server.llm_router is not None, (
"llm_router is None after client_no_auth fixture initialized. "
"This indicates a router initialization issue that should be investigated."
)
try:
with mock.patch.object(

View file

@ -427,16 +427,14 @@ async def test_acompletion_with_mcp_adds_metadata_to_streaming(monkeypatch):
assert len(all_chunks) > 0
# Verify mcp_list_tools is in the first chunk
first_chunk = all_chunks[0] if all_chunks else None
assert first_chunk is not None, "Should have a first chunk"
if hasattr(first_chunk, "choices") and first_chunk.choices:
choice = first_chunk.choices[0]
if hasattr(choice, "delta") and choice.delta:
provider_fields = getattr(choice.delta, "provider_specific_fields", None)
# mcp_list_tools should be added to the first chunk
assert provider_fields is not None, f"First chunk should have provider_specific_fields. Delta: {choice.delta}"
assert "mcp_list_tools" in provider_fields, f"First chunk should have mcp_list_tools. Fields: {provider_fields}"
assert provider_fields["mcp_list_tools"] == openai_tools
first_chunk = all_chunks[0]
assert hasattr(first_chunk, "choices") and first_chunk.choices, "First chunk must have choices"
choice = first_chunk.choices[0]
assert hasattr(choice, "delta") and choice.delta, "First choice must have delta"
provider_fields = getattr(choice.delta, "provider_specific_fields", None)
assert provider_fields is not None, f"First chunk should have provider_specific_fields. Delta: {choice.delta}"
assert "mcp_list_tools" in provider_fields, f"First chunk should have mcp_list_tools. Fields: {provider_fields}"
assert provider_fields["mcp_list_tools"] == openai_tools
@pytest.mark.asyncio
@ -625,7 +623,7 @@ async def test_acompletion_with_mcp_streaming_metadata_in_correct_chunks(monkeyp
],
), # Final chunk with tool_calls
]
follow_up_chunks = [
create_chunk("Hello"),
create_chunk(" world", finish_reason="stop"),
@ -785,21 +783,21 @@ async def test_acompletion_with_mcp_streaming_metadata_in_correct_chunks(monkeyp
assert initial_final_chunk is not None, "Should have a final chunk from initial response"
# Verify mcp_list_tools is in the first chunk
if hasattr(first_chunk, "choices") and first_chunk.choices:
choice = first_chunk.choices[0]
if hasattr(choice, "delta") and choice.delta:
provider_fields = getattr(choice.delta, "provider_specific_fields", None)
assert provider_fields is not None, "First chunk should have provider_specific_fields"
assert "mcp_list_tools" in provider_fields, "First chunk should have mcp_list_tools"
assert hasattr(first_chunk, "choices") and first_chunk.choices, "First chunk must have choices"
first_choice = first_chunk.choices[0]
assert hasattr(first_choice, "delta") and first_choice.delta, "First choice must have delta"
first_provider_fields = getattr(first_choice.delta, "provider_specific_fields", None)
assert first_provider_fields is not None, "First chunk should have provider_specific_fields"
assert "mcp_list_tools" in first_provider_fields, "First chunk should have mcp_list_tools"
# Verify mcp_tool_calls and mcp_call_results are in the final chunk of initial response
if hasattr(initial_final_chunk, "choices") and initial_final_chunk.choices:
choice = initial_final_chunk.choices[0]
if hasattr(choice, "delta") and choice.delta:
provider_fields = getattr(choice.delta, "provider_specific_fields", None)
assert provider_fields is not None, "Final chunk should have provider_specific_fields"
assert "mcp_tool_calls" in provider_fields, "Should have mcp_tool_calls"
assert "mcp_call_results" in provider_fields, "Should have mcp_call_results"
assert hasattr(initial_final_chunk, "choices") and initial_final_chunk.choices, "Final chunk must have choices"
final_choice = initial_final_chunk.choices[0]
assert hasattr(final_choice, "delta") and final_choice.delta, "Final choice must have delta"
final_provider_fields = getattr(final_choice.delta, "provider_specific_fields", None)
assert final_provider_fields is not None, "Final chunk should have provider_specific_fields"
assert "mcp_tool_calls" in final_provider_fields, "Should have mcp_tool_calls"
assert "mcp_call_results" in final_provider_fields, "Should have mcp_call_results"
@pytest.mark.asyncio