diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index de5b3b32105..a2e18a62638 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -22,78 +22,6 @@ def redis_no_ping(): yield -@pytest.mark.parametrize("namespace", [None, "test"]) -@pytest.mark.asyncio -async def test_redis_cache_async_increment(namespace, monkeypatch, redis_no_ping): - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache(namespace=namespace) - # Create an AsyncMock for the Redis client - mock_redis_instance = AsyncMock() - - # Make sure the mock can be used as an async context manager - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - - assert redis_cache is not None - - expected_key = "test:test" if namespace else "test" - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - # Call async_set_cache - await redis_cache.async_increment(key=expected_key, value=1) - - # Verify that the set method was called on the mock Redis instance - mock_redis_instance.incrbyfloat.assert_called_once_with( - name=expected_key, amount=1 - ) - - -@pytest.mark.asyncio -async def test_redis_cache_async_increment_refresh_ttl_true_bumps_existing_ttl( - monkeypatch, redis_no_ping -): - """With refresh_ttl=True, every increment should call expire() to bump - the TTL, even when the key already has a TTL (counter-style use).""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - mock_redis_instance = AsyncMock() - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - mock_redis_instance.ttl.return_value = 42 # key already has ~42s left - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - await redis_cache.async_increment( - key="spend:team_member:u:t", value=0.05, refresh_ttl=True - ) - - mock_redis_instance.expire.assert_awaited_once_with("spend:team_member:u:t", 60) - - -@pytest.mark.asyncio -async def test_redis_cache_async_increment_default_does_not_bump_existing_ttl( - monkeypatch, redis_no_ping -): - """Default (refresh_ttl=False) preserves window-style semantics: TTL is - set only on first creation, never refreshed (used by rate-limit windows).""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - mock_redis_instance = AsyncMock() - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - mock_redis_instance.ttl.return_value = 42 # key already has ~42s left - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - await redis_cache.async_increment(key="rate_limit:window", value=1) - - mock_redis_instance.expire.assert_not_awaited() - - @pytest.mark.parametrize("namespace", [None, "litellm"]) @pytest.mark.asyncio async def test_async_delete_cache_applies_namespace( @@ -140,42 +68,6 @@ async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping) assert client.connection_pool.connection_kwargs["socket_timeout"] == 1.0 -@pytest.mark.asyncio -async def test_redis_cache_async_batch_get_cache(monkeypatch, redis_no_ping): - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - # Create an AsyncMock for the Redis client - mock_redis_instance = AsyncMock() - - # Make sure the mock can be used as an async context manager - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - - # Setup the return value for mget - mock_redis_instance.mget.return_value = [ - b'{"key1": "value1"}', - None, - b'{"key3": "value3"}', - ] - - test_keys = ["key1", "key2", "key3"] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - # Call async_batch_get_cache - result = await redis_cache.async_batch_get_cache(key_list=test_keys) - - # Verify mget was called with the correct keys - mock_redis_instance.mget.assert_called_once() - - # Check that results were properly decoded - assert result["key1"] == {"key1": "value1"} - assert result["key2"] is None - assert result["key3"] == {"key3": "value3"} - - @pytest.mark.asyncio async def test_handle_lpop_count_for_older_redis_versions(monkeypatch): """Test the helper method that handles LPOP with count for Redis versions < 7.0""" @@ -202,41 +94,6 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch): assert mock_pipeline.execute.call_count == 2 -@pytest.mark.asyncio -async def test_async_rpush_pipeline_executes_all_operations(monkeypatch, redis_no_ping): - """Verify that multiple rpush ops are batched into a single pipeline execute""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.rpush = MagicMock() - mock_pipeline.execute = AsyncMock(return_value=[3, 5, 1]) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineRpushOperation - - rpush_list = [ - RedisPipelineRpushOperation(key="key1", values=["a", "b"]), - RedisPipelineRpushOperation(key="key2", values=["c"]), - RedisPipelineRpushOperation(key="key3", values=["d", "e", "f"]), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - result = await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) - - assert result == [3, 5, 1] - assert mock_pipeline.rpush.call_count == 3 - mock_pipeline.rpush.assert_any_call("key1", "a", "b") - mock_pipeline.rpush.assert_any_call("key2", "c") - mock_pipeline.rpush.assert_any_call("key3", "d", "e", "f") - mock_pipeline.execute.assert_called_once() - - @pytest.mark.asyncio async def test_async_rpush_pipeline_empty_list_returns_empty( monkeypatch, redis_no_ping @@ -256,183 +113,6 @@ async def test_async_rpush_pipeline_empty_list_returns_empty( mock_redis_instance.pipeline.assert_not_called() -@pytest.mark.asyncio -async def test_async_rpush_pipeline_raises_on_redis_error(monkeypatch, redis_no_ping): - """Pipeline errors should propagate""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.rpush = MagicMock() - mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down")) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineRpushOperation - - rpush_list = [RedisPipelineRpushOperation(key="key1", values=["a"])] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(ConnectionError, match="Redis down"): - await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) - - -@pytest.mark.asyncio -async def test_async_lpop_pipeline_single_round_trip(monkeypatch, redis_no_ping): - """Verify that multiple lpop ops are batched into a single pipeline execute""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "7.0.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - mock_pipeline.execute = AsyncMock( - return_value=[ - [b"val1", b"val2"], # key1 results - None, # key2 empty - [b"val3"], # key3 results - ] - ) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [ - RedisPipelineLpopOperation(key="key1", count=10), - RedisPipelineLpopOperation(key="key2", count=10), - RedisPipelineLpopOperation(key="key3", count=5), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - assert len(results) == 3 - assert results[0] == ["val1", "val2"] - assert results[1] is None - assert results[2] == ["val3"] - mock_pipeline.execute.assert_called_once() - - -@pytest.mark.asyncio -async def test_async_lpop_pipeline_redis_lt7_regroups_flat_results( - monkeypatch, redis_no_ping -): - """Verify Redis < 7 fallback issues individual LPOPs and regroups correctly""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "6.2.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - - # With count=3 for key1 and count=2 for key2, we get 5 individual LPOP commands - # Simulate: key1 has 2 values then None, key2 has 1 value then None - mock_pipeline.execute = AsyncMock( - return_value=[ - b"val1", - b"val2", - None, # 3 LPOPs for key1 - b"val3", - None, # 2 LPOPs for key2 - ] - ) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [ - RedisPipelineLpopOperation(key="key1", count=3), - RedisPipelineLpopOperation(key="key2", count=2), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - assert len(results) == 2 - assert results[0] == ["val1", "val2"] # 2 values, None filtered out - assert results[1] == ["val3"] # 1 value, None filtered out - # All 5 individual LPOPs should be queued, but only 1 execute() call - assert mock_pipeline.lpop.call_count == 5 - mock_pipeline.execute.assert_called_once() - - -@pytest.mark.asyncio -async def test_async_rpush_pipeline_raises_on_per_command_error( - monkeypatch, redis_no_ping -): - """Verify that per-command errors in pipeline results are raised, not silently dropped""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.rpush = MagicMock() - # Simulate: first RPUSH succeeds, second returns a per-command error - mock_pipeline.execute = AsyncMock(return_value=[3, Exception("WRONGTYPE")]) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineRpushOperation - - rpush_list = [ - RedisPipelineRpushOperation(key="key1", values=["a"]), - RedisPipelineRpushOperation(key="key2", values=["b"]), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(Exception, match="WRONGTYPE"): - await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) - - -@pytest.mark.asyncio -async def test_async_lpop_pipeline_raises_on_per_command_error( - monkeypatch, redis_no_ping -): - """Verify that per-command errors in LPOP pipeline results are raised, not silently dropped""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "7.0.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - # Simulate: first LPOP succeeds, second returns a per-command error - mock_pipeline.execute = AsyncMock(return_value=[[b"val1"], Exception("WRONGTYPE")]) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [ - RedisPipelineLpopOperation(key="key1", count=10), - RedisPipelineLpopOperation(key="key2", count=10), - ] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(Exception, match="WRONGTYPE"): - await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - @pytest.mark.asyncio async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping): """Empty lpop_list should return empty list without touching Redis""" @@ -450,111 +130,6 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping): mock_redis_instance.pipeline.assert_not_called() -@pytest.mark.asyncio -async def test_async_lpop_pipeline_propagates_redis_exception( - monkeypatch, redis_no_ping -): - """Pipeline errors should propagate""" - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - redis_cache = RedisCache() - redis_cache.redis_version = "7.0.0" - - mock_redis_instance = AsyncMock() - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - mock_pipeline.lpop = MagicMock() - mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down")) - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - from litellm.types.caching import RedisPipelineLpopOperation - - lpop_list = [RedisPipelineLpopOperation(key="key1", count=10)] - - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - with pytest.raises(ConnectionError, match="Redis down"): - await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "redis_version", - [ - # Standard cases - "7.0.0", # Standard Redis string version - 7.0, # Valkey/ElastiCache float version (THE BUG this fix addresses) - 7, # Integer version (e.g., from some Redis forks) - # Version < 7 - "6", # String without dots, version < 7 - # Malformed versions (fallback to 7) - "latest", # Non-numeric version - "", # Empty string - -7.0, # Negative float - # Format variations - " 7.0.0 ", # Whitespace (should be stripped) - "7.0.0-rc1", # Version with suffix - "10.0.0", # Double digit major version - ], -) -async def test_async_lpop_with_float_redis_version( - monkeypatch, redis_no_ping, redis_version -): - """ - Test async_lpop with various Redis version formats (especially float). - - This test specifically addresses the issue where AWS ElastiCache Valkey - returns redis_version as a float (e.g., 7.0) instead of a string (e.g., "7.0.0"), - which caused a 'float' object has no attribute 'split' error when trying to - use the Redis transaction buffer feature. - - The fix converts the version to a string and handles edge cases like: - - Floats (7.0) and integers (7) - - Strings with/without dots ("7" vs "7.0.0") - - Malformed versions ("v7.0.0", "latest") - fallback to version 7 - - Whitespace (" 7.0.0 ") - - Negative versions (fallback to version 7) - - Related: Database deadlock issues when use_redis_transaction_buffer is enabled. - """ - monkeypatch.setenv("REDIS_HOST", "https://my-test-host") - - # Create RedisCache instance - redis_cache = RedisCache() - redis_cache.redis_version = redis_version # Set the version to test - - # Create an AsyncMock for the Redis client - mock_redis_instance = AsyncMock() - mock_redis_instance.__aenter__.return_value = mock_redis_instance - mock_redis_instance.__aexit__.return_value = None - - # Mock lpop to return a test value (Redis >= 7.0 behavior) - mock_redis_instance.lpop.return_value = [b"value1", b"value2"] - - # Mock pipeline for Redis < 7.0 (used when major_version < 7) - mock_pipeline = MagicMock() - mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) - mock_pipeline.__aexit__ = AsyncMock(return_value=None) - # Make pipeline() a regular method (not async) that returns the mock - mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) - - # Mock handle_lpop_count_for_older_redis_versions for Redis < 7 - with patch.object( - redis_cache, - "handle_lpop_count_for_older_redis_versions", - return_value=[b"value1", b"value2"], - ): - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): - # Call async_lpop with count - this should not raise AttributeError - result = await redis_cache.async_lpop(key="test_key", count=2) - - # Verify the method completed without error - assert result is not None - - # LIT-3374: the namespace must be applied uniformly across every key-taking # Redis operation, not just get/set/increment. Before the fix these paths wrote # or read raw keys, so with a namespace configured the prefixed keys other 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 7bd1d7a6031..87d67e0e8b7 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -181,28 +181,6 @@ async def test_force_ipv4_transport(): litellm.disable_aiohttp_transport = original_disable -@pytest.mark.asyncio -async def test_ssl_context_transport(): - """Test transport creation with SSL context""" - # Create a test SSL context - ssl_context = ssl.create_default_context() - - transport = AsyncHTTPHandler._create_async_transport(ssl_context=ssl_context) - assert transport is not None - - try: - if isinstance(transport, LiteLLMAiohttpTransport): - # Get the client session and verify SSL context is passed through - client_session = transport._get_valid_client_session() - assert isinstance(client_session, ClientSession) - assert isinstance(client_session.connector, TCPConnector) - # Verify the connector has SSL context set by checking if it's using SSL - assert client_session.connector._ssl is not None - finally: - if isinstance(transport, LiteLLMAiohttpTransport): - await transport.aclose() - - @pytest.mark.asyncio async def test_aiohttp_disabled_transport(): """Test transport creation with aiohttp disabled""" @@ -339,44 +317,6 @@ async def test_ssl_context_with_shared_session(): litellm.disable_aiohttp_transport = original_disable -@pytest.mark.asyncio -async def test_aiohttp_transport_trust_env_setting(monkeypatch): - """Test that trust_env setting is properly configured in aiohttp transport""" - transports = [] - try: - # Test 1: Default trust_env behavior - transport = AsyncHTTPHandler._create_aiohttp_transport() - transports.append(transport) - 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) - assert client_session._trust_env == default_trust_env - - # Test 2: Environment variable override - monkeypatch.setenv("AIOHTTP_TRUST_ENV", "True") - transport_with_env = AsyncHTTPHandler._create_aiohttp_transport() - transports.append(transport_with_env) - client_session_with_env = transport_with_env._get_valid_client_session() - - # Should be True when environment variable is set - assert client_session_with_env._trust_env is True - - # Test 3: Verify environment variable with False value - 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() - ) - - # Should respect the litellm.aiohttp_trust_env setting when env var is False - assert client_session_with_false_env._trust_env == default_trust_env - finally: - for t in transports: - await t.aclose() - - 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.""" @@ -443,36 +383,6 @@ async def test_create_aiohttp_transport_with_shared_session(): assert not callable(transport.client) # Should not be callable -@pytest.mark.asyncio -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 - - -@pytest.mark.asyncio -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 - - @pytest.mark.asyncio async def test_async_handler_with_shared_session(): """Test AsyncHTTPHandler initialization with shared session""" @@ -622,27 +532,6 @@ async def test_session_reuse_integration(): await client2.close() -@pytest.mark.asyncio -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 - assert transport3.client is mock_valid_session # Should reuse session - - @pytest.mark.parametrize( "env_curve,litellm_curve,expected_curve,should_call", [ diff --git a/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py b/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py index f8bd83fc7df..1043c26c6ec 100644 --- a/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py +++ b/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py @@ -1,15 +1,10 @@ """ -Tests for LlmPassthroughRouteHandler and the guardrail_translation_mappings registry. +Tests for the guardrail_translation_mappings registry. Validates: - allm_passthrough_route is registered in the mappings (regression: this was the bug) -- Bedrock provider is dispatched to BedrockPassthroughGuardrailHandler -- Unknown provider skips apply_guardrail """ -import pytest -from unittest.mock import AsyncMock, MagicMock, patch - from litellm.llms.pass_through.guardrail_translation import ( guardrail_translation_mappings, ) @@ -40,185 +35,3 @@ class TestRegistry: is PassThroughEndpointHandler ) - -def _make_guardrail() -> MagicMock: - g = MagicMock() - g.guardrail_name = "test-guard" - g.apply_guardrail = AsyncMock(return_value={"texts": []}) - g.skip_system_message_in_guardrail = False - g.skip_tool_message_in_guardrail = False - return g - - -class TestLlmPassthroughRouteHandlerInput: - @pytest.mark.asyncio - async def test_bedrock_provider_delegates_to_bedrock_handler(self): - handler = LlmPassthroughRouteHandler() - data = { - "custom_llm_provider": "bedrock", - "endpoint": "model/anthropic.claude-3-sonnet/converse", - "data": {"messages": [{"role": "user", "content": [{"text": "hi"}]}]}, - } - guardrail = _make_guardrail() - - await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) - - guardrail.apply_guardrail.assert_called_once() - - @pytest.mark.asyncio - async def test_unknown_provider_skips_apply_guardrail(self): - handler = LlmPassthroughRouteHandler() - data = { - "custom_llm_provider": "some_unknown_provider", - "endpoint": "v1/chat/completions", - "data": {"messages": [{"role": "user", "content": "hi"}]}, - } - guardrail = _make_guardrail() - - result = await handler.process_input_messages( - data=data, guardrail_to_apply=guardrail - ) - - guardrail.apply_guardrail.assert_not_called() - assert result is data - - @pytest.mark.asyncio - async def test_missing_provider_skips(self): - handler = LlmPassthroughRouteHandler() - data = {"endpoint": "foo/bar", "data": {}} - guardrail = _make_guardrail() - - result = await handler.process_input_messages( - data=data, guardrail_to_apply=guardrail - ) - - guardrail.apply_guardrail.assert_not_called() - assert result is data - - -class TestLlmPassthroughRouteHandlerOutput: - @pytest.mark.asyncio - async def test_bedrock_provider_delegates_output_to_bedrock_handler(self): - handler = LlmPassthroughRouteHandler() - response = { - "output": { - "message": { - "role": "assistant", - "content": [{"text": "hello"}], - } - } - } - request_data = { - "custom_llm_provider": "bedrock", - "endpoint": "model/anthropic.claude-3-sonnet/converse", - } - guardrail = _make_guardrail() - - await handler.process_output_response( - response=response, - guardrail_to_apply=guardrail, - request_data=request_data, - ) - - guardrail.apply_guardrail.assert_called_once() - - @pytest.mark.asyncio - async def test_unknown_provider_skips_output(self): - handler = LlmPassthroughRouteHandler() - response = {"some": "response"} - request_data = {"custom_llm_provider": "unknown"} - guardrail = _make_guardrail() - - result = await handler.process_output_response( - response=response, - guardrail_to_apply=guardrail, - request_data=request_data, - ) - - guardrail.apply_guardrail.assert_not_called() - assert result is response - - -class TestDeAnonymizeEventStream: - @pytest.mark.asyncio - async def test_bedrock_provider_dispatches_to_handler(self): - body = b"original-stream-bytes" - expected = b"de-anonymized-bytes" - proxy_logging_obj = MagicMock() - user_api_key_dict = MagicMock() - - with patch( - "litellm.llms.bedrock.passthrough.guardrail_translation.handler." - "BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", - new=AsyncMock(return_value=expected), - ) as mock_handler: - result = await LlmPassthroughRouteHandler.de_anonymize_event_stream( - body_bytes=body, - proxy_logging_obj=proxy_logging_obj, - user_api_key_dict=user_api_key_dict, - data={"custom_llm_provider": "bedrock"}, - ) - - mock_handler.assert_awaited_once() - assert result == expected - - @pytest.mark.asyncio - async def test_unknown_provider_returns_original_bytes(self): - body = b"original-stream-bytes" - - result = await LlmPassthroughRouteHandler.de_anonymize_event_stream( - body_bytes=body, - proxy_logging_obj=MagicMock(), - user_api_key_dict=MagicMock(), - data={"custom_llm_provider": "anthropic"}, - ) - - assert result is body - - @pytest.mark.asyncio - async def test_missing_provider_returns_original_bytes(self): - body = b"original-stream-bytes" - - result = await LlmPassthroughRouteHandler.de_anonymize_event_stream( - body_bytes=body, - proxy_logging_obj=MagicMock(), - user_api_key_dict=MagicMock(), - data={}, - ) - - assert result is body - - -class TestSupportsEventStreamDeAnonymization: - def test_bedrock_converse_stream_is_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream" - ) - is True - ) - - def test_bedrock_invoke_stream_is_not_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "bedrock", - "model/us.amazon.nova-lite-v1:0/invoke-with-response-stream", - ) - is False - ) - - def test_unknown_provider_is_not_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - "anthropic", "model/foo/converse-stream" - ) - is False - ) - - def test_missing_provider_is_not_supported(self): - assert ( - LlmPassthroughRouteHandler.supports_event_stream_de_anonymization( - None, "model/foo/converse-stream" - ) - is False - ) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index bad76864ca7..087aaec9215 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1504,50 +1504,6 @@ def test_team_info_masking(): assert "public-test-key" not in str(exc_info.value) -def test_embedding_input_array_of_tokens(client_no_auth): - """ - Test to bypass decoding input as array of tokens for selected providers - - Ref: https://github.com/BerriAI/litellm/issues/10113 - """ - from litellm.proxy import proxy_server - - # 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( - proxy_server.llm_router, - "aembedding", - return_value=example_embedding_result, - ) as mock_aembedding: - test_data = { - "model": "vllm_embed_model", - "input": [[2046, 13269, 158208]], - } - - response = client_no_auth.post("/v1/embeddings", json=test_data) - - # Assert that aembedding was called, and that input was not modified - mock_aembedding.assert_called_once() - call_args, call_kwargs = mock_aembedding.call_args - assert call_kwargs["model"] == "vllm_embed_model" - assert call_kwargs["input"] == [[2046, 13269, 158208]] - - assert response.status_code == 200 - result = response.json() - print(len(result["data"][0]["embedding"])) - assert ( - len(result["data"][0]["embedding"]) > 10 - ) # this usually has len==1536 so - except Exception as e: - pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}") - - @pytest.mark.asyncio async def test_get_all_team_models(): """