mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #34475 from BerriAI/litellm_/test-coverage-mutation-analysis-e42223
test: remove tests that mutation analysis proved assert nothing
This commit is contained in:
commit
7047a37f2f
4 changed files with 1 additions and 768 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue