mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #20595 from jquinter/fix/test-parallelization-isolation
fix: improve test isolation for parallel execution
This commit is contained in:
commit
cca4a8699a
6 changed files with 67 additions and 44 deletions
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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')
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue