diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index a7b7187c8ff..f98550b9909 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -13,7 +13,10 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm from litellm.llms.custom_httpx.aiohttp_transport import LiteLLMAiohttpTransport -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_ssl_configuration +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + get_ssl_configuration, +) @pytest.mark.asyncio @@ -21,7 +24,7 @@ async def test_ssl_security_level(monkeypatch): # Ensure aiohttp transport is enabled for this test original_disable = litellm.disable_aiohttp_transport litellm.disable_aiohttp_transport = False - + try: with patch.dict(os.environ, clear=True): # Set environment variable for SSL security level @@ -125,7 +128,7 @@ async def test_ssl_verification_with_aiohttp_transport(): # Ensure aiohttp transport is enabled for this test original_disable = litellm.disable_aiohttp_transport litellm.disable_aiohttp_transport = False - + try: litellm_async_client = AsyncHTTPHandler(ssl_verify=False) @@ -241,7 +244,7 @@ async def test_aiohttp_transport_trust_env_setting(monkeypatch): client_session = transport._get_valid_client_session() # Default should be False (litellm.aiohttp_trust_env default) - default_trust_env = getattr(litellm, 'aiohttp_trust_env', False) + default_trust_env = getattr(litellm, "aiohttp_trust_env", False) assert client_session._trust_env == default_trust_env # Test 2: Environment variable override @@ -257,7 +260,9 @@ async def test_aiohttp_transport_trust_env_setting(monkeypatch): monkeypatch.setenv("AIOHTTP_TRUST_ENV", "False") transport_with_false_env = AsyncHTTPHandler._create_aiohttp_transport() transports.append(transport_with_false_env) - client_session_with_false_env = transport_with_false_env._get_valid_client_session() + client_session_with_false_env = ( + transport_with_false_env._get_valid_client_session() + ) # Should respect the litellm.aiohttp_trust_env setting when env var is False assert client_session_with_false_env._trust_env == default_trust_env @@ -270,25 +275,25 @@ def test_get_ssl_configuration(): """Test that get_ssl_configuration() returns a proper SSL context with certifi CA bundle when no environment variables are set.""" from litellm.llms.custom_httpx.http_handler import _ssl_context_cache - + # Clear cache to ensure ssl.create_default_context is called _ssl_context_cache.clear() - + with patch.dict(os.environ, clear=True): - with patch('ssl.create_default_context') as mock_create_context: + with patch("ssl.create_default_context") as mock_create_context: # Mock the return value mock_ssl_context = MagicMock(spec=ssl.SSLContext) mock_ssl_context.set_ciphers = MagicMock() mock_ssl_context.minimum_version = ssl.TLSVersion.TLSv1_2 mock_create_context.return_value = mock_ssl_context - + # Call the static method result = get_ssl_configuration() - + # Verify ssl.create_default_context was called with certifi's CA file expected_ca_file = certifi.where() mock_create_context.assert_called_once_with(cafile=expected_ca_file) - + # Verify it returns the mocked SSL context assert result == mock_ssl_context @@ -297,10 +302,10 @@ def test_get_ssl_configuration_integration(): """Integration test that _get_ssl_context() returns a working SSL context""" # Call the static method without mocking ssl_context = get_ssl_configuration() - + # Verify it returns an SSLContext instance assert isinstance(ssl_context, ssl.SSLContext) - + # Verify it has basic SSL context properties assert ssl_context.protocol is not None assert ssl_context.verify_mode is not None @@ -309,22 +314,24 @@ def test_get_ssl_configuration_integration(): # Session Reuse Tests class MockClientSession: """Mock ClientSession that is not callable""" + def __init__(self): self.closed = False + @pytest.mark.asyncio async def test_create_aiohttp_transport_with_shared_session(): """Test that _create_aiohttp_transport reuses shared session when provided""" from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - + # Create a mock shared session that's not callable mock_session = MockClientSession() - + # Test with shared session transport = AsyncHTTPHandler._create_aiohttp_transport( shared_session=mock_session # type: ignore ) - + # Verify the transport uses the shared session directly assert transport.client is mock_session assert not callable(transport.client) # Should not be callable @@ -334,10 +341,10 @@ async def test_create_aiohttp_transport_with_shared_session(): async def test_create_aiohttp_transport_without_shared_session(): """Test that _create_aiohttp_transport creates new session when none provided""" from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - + # Test without shared session transport = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None) - + # Verify the transport uses a lambda function (for backward compatibility) assert callable(transport.client) # Should be a lambda function @@ -346,16 +353,16 @@ async def test_create_aiohttp_transport_without_shared_session(): async def test_create_aiohttp_transport_with_closed_session(): """Test that _create_aiohttp_transport creates new session when shared session is closed""" from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - + # Create a mock closed session mock_session = MockClientSession() mock_session.closed = True - + # Test with closed session transport = AsyncHTTPHandler._create_aiohttp_transport( shared_session=mock_session # type: ignore ) - + # Verify the transport creates a new session (lambda function) assert callable(transport.client) # Should be a lambda function @@ -364,13 +371,13 @@ async def test_create_aiohttp_transport_with_closed_session(): async def test_async_handler_with_shared_session(): """Test AsyncHTTPHandler initialization with shared session""" from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - + # Create a mock shared session mock_session = MockClientSession() - + # Create handler with shared session handler = AsyncHTTPHandler(shared_session=mock_session) # type: ignore - + # Verify the handler was created successfully assert handler is not None assert handler.client is not None @@ -379,7 +386,10 @@ async def test_async_handler_with_shared_session(): @pytest.mark.asyncio async def test_get_async_httpx_client_with_shared_session(): """Test get_async_httpx_client with shared session""" - from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, AsyncHTTPHandler as AsyncHTTPHandlerReload + from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + AsyncHTTPHandler as AsyncHTTPHandlerReload, + ) from litellm.types.utils import LlmProviders # Create a mock shared session @@ -387,8 +397,7 @@ async def test_get_async_httpx_client_with_shared_session(): # Test with shared session client = get_async_httpx_client( - llm_provider=LlmProviders.ANTHROPIC, - shared_session=mock_session # type: ignore + llm_provider=LlmProviders.ANTHROPIC, shared_session=mock_session # type: ignore ) # Verify the client was created successfully @@ -400,13 +409,15 @@ async def test_get_async_httpx_client_with_shared_session(): @pytest.mark.asyncio async def test_get_async_httpx_client_without_shared_session(): """Test get_async_httpx_client without shared session (backward compatibility)""" - from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, AsyncHTTPHandler as AsyncHTTPHandlerReload + from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + AsyncHTTPHandler as AsyncHTTPHandlerReload, + ) from litellm.types.utils import LlmProviders # Test without shared session client = get_async_httpx_client( - llm_provider=LlmProviders.ANTHROPIC, - shared_session=None + llm_provider=LlmProviders.ANTHROPIC, shared_session=None ) # Verify the client was created successfully @@ -419,18 +430,18 @@ async def test_get_async_httpx_client_without_shared_session(): async def test_session_reuse_chain(): """Test that session is properly passed through the entire call chain""" from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - + # Create a mock shared session mock_session = MockClientSession() - + # Test the entire chain transport = AsyncHTTPHandler._create_async_transport( shared_session=mock_session # type: ignore ) - + # Verify the transport was created assert transport is not None - + # Test AsyncHTTPHandler creation handler = AsyncHTTPHandler(shared_session=mock_session) # type: ignore assert handler is not None @@ -440,40 +451,43 @@ def test_shared_session_parameter_in_acompletion(): """Test that acompletion function accepts shared_session parameter""" import inspect from litellm.main import acompletion - + # Get the function signature sig = inspect.signature(acompletion) params = list(sig.parameters.keys()) - + # Verify shared_session parameter exists - assert 'shared_session' in params - + assert "shared_session" in params + # Verify the parameter type annotation - shared_session_param = sig.parameters['shared_session'] - assert 'ClientSession' in str(shared_session_param.annotation) + shared_session_param = sig.parameters["shared_session"] + assert "ClientSession" in str(shared_session_param.annotation) def test_shared_session_parameter_in_completion(): """Test that completion function accepts shared_session parameter""" import inspect from litellm.main import completion - + # Get the function signature sig = inspect.signature(completion) params = list(sig.parameters.keys()) - + # Verify shared_session parameter exists - assert 'shared_session' in params - + assert "shared_session" in params + # Verify the parameter type annotation - shared_session_param = sig.parameters['shared_session'] - assert 'ClientSession' in str(shared_session_param.annotation) + shared_session_param = sig.parameters["shared_session"] + assert "ClientSession" in str(shared_session_param.annotation) @pytest.mark.asyncio async def test_session_reuse_integration(): """Integration test for session reuse functionality""" - from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, AsyncHTTPHandler as AsyncHTTPHandlerReload + from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + AsyncHTTPHandler as AsyncHTTPHandlerReload, + ) from litellm.types.utils import LlmProviders # Create a mock session @@ -481,13 +495,11 @@ async def test_session_reuse_integration(): # Create two clients with the same session client1 = get_async_httpx_client( - llm_provider=LlmProviders.ANTHROPIC, - shared_session=mock_session # type: ignore + llm_provider=LlmProviders.ANTHROPIC, shared_session=mock_session # type: ignore ) client2 = get_async_httpx_client( - llm_provider=LlmProviders.OPENAI, - shared_session=mock_session # type: ignore + llm_provider=LlmProviders.OPENAI, shared_session=mock_session # type: ignore ) # Both clients should be created successfully @@ -508,17 +520,17 @@ async def test_session_reuse_integration(): async def test_session_validation(): """Test that session validation works correctly""" from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - + # Test with None session transport1 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None) assert callable(transport1.client) # Should create lambda - + # Test with closed session mock_closed_session = MockClientSession() mock_closed_session.closed = True transport2 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_closed_session) # type: ignore assert callable(transport2.client) # Should create lambda - + # Test with valid session mock_valid_session = MockClientSession() transport3 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_valid_session) # type: ignore @@ -530,41 +542,42 @@ async def test_session_validation(): [ # env_curve: SSL_ECDH_CURVE env var | litellm_curve: litellm.ssl_ecdh_curve variable # expected_curve: curve that should be set | should_call: whether set_ecdh_curve() should be called - # Valid configurations - ("X25519", None, "X25519", True), # Env var only - ("prime256v1", None, "prime256v1", True), # Different valid curve - (None, "secp384r1", "secp384r1", True), # litellm variable only - ("X25519", "secp521r1", "X25519", True), # Env var takes precedence + ("X25519", None, "X25519", True), # Env var only + ("prime256v1", None, "prime256v1", True), # Different valid curve + (None, "secp384r1", "secp384r1", True), # litellm variable only + ("X25519", "secp521r1", "X25519", True), # Env var takes precedence # Empty/None configurations - should skip - ("", None, None, False), # Empty string - skip configuration - (None, None, None, False), # None value - skip configuration - ] + ("", None, None, False), # Empty string - skip configuration + (None, None, None, False), # None value - skip configuration + ], ) -def test_ssl_ecdh_curve(env_curve, litellm_curve, expected_curve, should_call, monkeypatch): +def test_ssl_ecdh_curve( + env_curve, litellm_curve, expected_curve, should_call, monkeypatch +): """Test SSL ECDH curve configuration with valid curves and precedence""" from litellm.llms.custom_httpx.http_handler import _ssl_context_cache - + # Clear cache to ensure fresh SSL context creation _ssl_context_cache.clear() - + with patch.dict(os.environ, clear=True): if env_curve: monkeypatch.setenv("SSL_ECDH_CURVE", env_curve) - + original_value = litellm.ssl_ecdh_curve try: litellm.ssl_ecdh_curve = litellm_curve - + # Create a real SSL context and patch set_ecdh_curve on it # We need a real SSLContext instance (not a MagicMock) because _create_ssl_context # calls methods like set_ciphers() and minimum_version that require a real context. # We patch set_ecdh_curve specifically to verify it's called with the correct curve. real_ssl_context = ssl.create_default_context() - with patch('ssl.create_default_context', return_value=real_ssl_context): - with patch.object(real_ssl_context, 'set_ecdh_curve') as mock_set_curve: + with patch("ssl.create_default_context", return_value=real_ssl_context): + with patch.object(real_ssl_context, "set_ecdh_curve") as mock_set_curve: ssl_context = get_ssl_configuration() - + if should_call: mock_set_curve.assert_called_once_with(expected_curve) else: