fix: FLAKY tests

This commit is contained in:
Ishaan Jaffer 2026-01-24 11:13:44 -08:00
parent 9cdd7a8fd2
commit bd38374a45
4 changed files with 56 additions and 43 deletions

View File

@ -52,16 +52,19 @@ class TestMCPClient:
self, mock_session, mock_stdio_client
):
"""Test successful stdio connection"""
# Setup mocks
# Setup mocks - create proper async context manager
mock_transport = (MagicMock(), MagicMock())
mock_stdio_client.return_value.__aenter__ = AsyncMock(
return_value=mock_transport
)
mock_stdio_ctx = AsyncMock()
mock_stdio_ctx.__aenter__.return_value = mock_transport
mock_stdio_ctx.__aexit__.return_value = None
mock_stdio_client.return_value = mock_stdio_ctx
mock_session_instance = MagicMock()
mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance)
mock_session_instance = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session.return_value = mock_session_instance
mock_session_ctx = AsyncMock()
mock_session_ctx.__aenter__.return_value = mock_session_instance
mock_session_ctx.__aexit__.return_value = None
mock_session.return_value = mock_session_ctx
stdio_config = MCPStdioConfig(
command="python", args=["-m", "my_mcp_server"], env={"DEBUG": "1"}
@ -94,22 +97,23 @@ class TestMCPClient:
self, mock_streamable_http_client
):
"""Test that MCP client uses SSL configuration from environment variables"""
# Setup mocks
# Setup mocks - create proper async context manager
mock_transport = (MagicMock(), MagicMock())
mock_streamable_http_client.return_value.__aenter__ = AsyncMock(
return_value=mock_transport
)
mock_http_ctx = AsyncMock()
mock_http_ctx.__aenter__.return_value = mock_transport
mock_http_ctx.__aexit__.return_value = None
mock_streamable_http_client.return_value = mock_http_ctx
# Mock the session
with patch(
"litellm.experimental_mcp_client.client.ClientSession"
) as mock_session:
mock_session_instance = MagicMock()
mock_session_instance.__aenter__ = AsyncMock(
return_value=mock_session_instance
)
mock_session_instance = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session.return_value = mock_session_instance
mock_session_ctx = AsyncMock()
mock_session_ctx.__aenter__.return_value = mock_session_instance
mock_session_ctx.__aexit__.return_value = None
mock_session.return_value = mock_session_ctx
client = MCPClient(
server_url="https://mcp-server.example.com",
@ -141,20 +145,23 @@ class TestMCPClient:
@patch.object(mcp_client_module, "sse_client")
async def test_mcp_client_ssl_verify_parameter(self, mock_sse_client):
"""Test that MCP client uses ssl_verify parameter when provided"""
# Setup mocks
# Setup mocks - create proper async context manager
mock_transport = (MagicMock(), MagicMock())
mock_sse_client.return_value.__aenter__ = AsyncMock(return_value=mock_transport)
mock_sse_ctx = AsyncMock()
mock_sse_ctx.__aenter__.return_value = mock_transport
mock_sse_ctx.__aexit__.return_value = None
mock_sse_client.return_value = mock_sse_ctx
# Mock the session
with patch(
"litellm.experimental_mcp_client.client.ClientSession"
) as mock_session:
mock_session_instance = MagicMock()
mock_session_instance.__aenter__ = AsyncMock(
return_value=mock_session_instance
)
mock_session_instance = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session.return_value = mock_session_instance
mock_session_ctx = AsyncMock()
mock_session_ctx.__aenter__.return_value = mock_session_instance
mock_session_ctx.__aexit__.return_value = None
mock_session.return_value = mock_session_ctx
# Test with ssl_verify=False
client = MCPClient(
@ -192,22 +199,23 @@ class TestMCPClient:
@patch.object(mcp_client_module, "streamable_http_client")
async def test_mcp_client_ssl_verify_custom_path(self, mock_streamable_http_client):
"""Test that MCP client uses custom CA bundle path from ssl_verify parameter"""
# Setup mocks
# Setup mocks - create proper async context manager
mock_transport = (MagicMock(), MagicMock())
mock_streamable_http_client.return_value.__aenter__ = AsyncMock(
return_value=mock_transport
)
mock_http_ctx = AsyncMock()
mock_http_ctx.__aenter__.return_value = mock_transport
mock_http_ctx.__aexit__.return_value = None
mock_streamable_http_client.return_value = mock_http_ctx
# Mock the session
with patch(
"litellm.experimental_mcp_client.client.ClientSession"
) as mock_session:
mock_session_instance = MagicMock()
mock_session_instance.__aenter__ = AsyncMock(
return_value=mock_session_instance
)
mock_session_instance = AsyncMock()
mock_session_instance.initialize = AsyncMock()
mock_session.return_value = mock_session_instance
mock_session_ctx = AsyncMock()
mock_session_ctx.__aenter__.return_value = mock_session_instance
mock_session_ctx.__aexit__.return_value = None
mock_session.return_value = mock_session_ctx
# Test with custom CA bundle path
custom_ca_path = "/custom/path/to/ca-bundle.pem"

View File

@ -1090,8 +1090,8 @@ async def test_google_generate_content_with_openai():
usage=mock_usage
)
# Use AsyncMock for proper async function mocking
with unittest.mock.patch("litellm.acompletion", new_callable=unittest.mock.AsyncMock) as mock_completion:
# Use AsyncMock for proper async function mocking - patch at the module level where it's imported
with unittest.mock.patch("litellm.google_genai.main.litellm.acompletion", new_callable=unittest.mock.AsyncMock) as mock_completion:
# Set the return value directly on the MagicMock
mock_completion.return_value = mock_response

View File

@ -25,7 +25,7 @@ def test_non_stream_response_when_stream_requested_sync():
the sync handler correctly transforms it to generate_content format.
"""
from litellm.types.utils import Choices
# Mock a non-stream response (ModelResponse with valid choices)
mock_response = ModelResponse(
id="test-123",
@ -70,7 +70,7 @@ async def test_non_stream_response_when_stream_requested_async():
the async handler correctly transforms it to generate_content format.
"""
from litellm.types.utils import Choices
# Mock a non-stream response (ModelResponse with valid choices)
mock_response = ModelResponse(
id="test-123",
@ -183,7 +183,8 @@ def test_stream_transformation_error_sync():
"translate_completion_output_params_streaming",
return_value=None
):
with patch("litellm.completion", return_value=mock_stream):
# Mock at the module level where it's imported
with patch("litellm.google_genai.adapters.handler.litellm.completion", return_value=mock_stream):
# Call the handler with stream=True and expect a ValueError
with pytest.raises(ValueError, match="Failed to transform streaming response"):
GenerateContentToCompletionHandler.generate_content_handler(
@ -209,7 +210,8 @@ async def test_stream_transformation_error_async():
"translate_completion_output_params_streaming",
return_value=None
):
with patch("litellm.acompletion", return_value=mock_stream):
# Mock at the module level where it's imported
with patch("litellm.google_genai.adapters.handler.litellm.acompletion", return_value=mock_stream):
# Call the handler with stream=True and expect a ValueError
with pytest.raises(ValueError, match="Failed to transform streaming response"):
await GenerateContentToCompletionHandler.async_generate_content_handler(
@ -225,11 +227,13 @@ def test_citation_metadata_transformation():
Test that citationMetadata.citationSources is properly transformed to citationMetadata.citations
to avoid Pydantic validation errors.
"""
from litellm.llms.gemini.google_genai.transformation import GoogleGenAIConfig
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from unittest.mock import MagicMock
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.gemini.google_genai.transformation import GoogleGenAIConfig
# Create a mock response with citationMetadata.citationSources (the problematic format)
mock_response_data = {
"candidates": [

View File

@ -15,7 +15,6 @@ import urllib.parse
from unittest.mock import MagicMock, patch
import litellm
from litellm import main as litellm_main
@ -473,6 +472,7 @@ async def test_extra_body_with_fallback(
@pytest.mark.parametrize("env_base", ["OPENAI_BASE_URL", "OPENAI_API_BASE"])
@pytest.mark.asyncio
@pytest.mark.flaky(retries=3, delay=1)
async def test_openai_env_base(
respx_mock: respx.MockRouter, env_base, openai_api_response, monkeypatch
):
@ -488,7 +488,8 @@ async def test_openai_env_base(
model = "gpt-4o"
messages = [{"role": "user", "content": "Hello, how are you?"}]
respx_mock.post(f"{expected_base_url}/chat/completions").respond(
# Ensure respx_mock is properly configured
respx_mock.route(host="localhost", port=12345).post("/v1/chat/completions").respond(
json={
"id": "chatcmpl-123",
"object": "chat.completion",