mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
761 lines
29 KiB
Python
761 lines
29 KiB
Python
import asyncio
|
|
import logging
|
|
import time
|
|
import uuid
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import pytest
|
|
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.caching.in_memory_cache import InMemoryCache
|
|
from litellm.caching.redis_cache import RedisCache, _redis_circuit_breaker_guard, _redis_circuit_breaker_guard_sync
|
|
from litellm.types.caching import RedisPipelineIncrementOperation
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_async_batch_get_cache_coalesces_concurrent_redis_reads():
|
|
dual_cache = DualCache(
|
|
redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10
|
|
)
|
|
keys = ["shared_a", "shared_b"]
|
|
start_gate = asyncio.Event()
|
|
|
|
async def _mock_async_batch_get_cache(key_list, parent_otel_span=None):
|
|
await asyncio.sleep(0.05)
|
|
return {k: None for k in key_list}
|
|
|
|
with patch.object(
|
|
dual_cache.redis_cache,
|
|
"async_batch_get_cache",
|
|
new=AsyncMock(side_effect=_mock_async_batch_get_cache),
|
|
) as mock_async_batch_get_cache:
|
|
|
|
async def worker():
|
|
await start_gate.wait()
|
|
return await dual_cache.async_batch_get_cache(keys=keys)
|
|
|
|
tasks = [asyncio.create_task(worker()) for _ in range(50)]
|
|
start_gate.set()
|
|
await asyncio.gather(*tasks)
|
|
|
|
assert mock_async_batch_get_cache.call_count == 1
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_async_batch_get_cache_rolls_back_redis_reservation_on_error():
|
|
dual_cache = DualCache(
|
|
redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10
|
|
)
|
|
keys = ["shared_a", "shared_b"]
|
|
|
|
with patch.object(
|
|
dual_cache.redis_cache,
|
|
"async_batch_get_cache",
|
|
new=AsyncMock(side_effect=RuntimeError("redis unavailable")),
|
|
) as mock_async_batch_get_cache:
|
|
first_result = await dual_cache.async_batch_get_cache(keys=keys)
|
|
second_result = await dual_cache.async_batch_get_cache(keys=keys)
|
|
|
|
assert first_result is None
|
|
assert second_result is None
|
|
assert mock_async_batch_get_cache.call_count == 2
|
|
assert "shared_a" not in dual_cache.last_redis_batch_access_time
|
|
assert "shared_b" not in dual_cache.last_redis_batch_access_time
|
|
|
|
|
|
def _redis_mock_for_sync_batch(redis_result: dict) -> MagicMock:
|
|
mock_redis = MagicMock(spec=RedisCache)
|
|
mock_redis.batch_get_cache.return_value = redis_result
|
|
return mock_redis
|
|
|
|
|
|
def _assert_sync_batch_used_blocking_client(dual_cache: DualCache, mock_redis: MagicMock) -> None:
|
|
with patch("asyncio.new_event_loop", side_effect=AssertionError("sync path must not create an event loop")):
|
|
result = dual_cache.batch_get_cache(keys=["lit6729_key"])
|
|
|
|
assert result == ["redis_value"]
|
|
mock_redis.batch_get_cache.assert_called_once_with(key_list=["lit6729_key"], parent_otel_span=None)
|
|
mock_redis.async_batch_get_cache.assert_not_called()
|
|
mock_redis.init_async_client.assert_not_called()
|
|
assert dual_cache.in_memory_cache.get_cache("lit6729_key") == "redis_value"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_batch_get_cache_uses_sync_redis_client_inside_running_loop():
|
|
"""
|
|
Regression test for LIT-6729: sync batch_get_cache ran async_batch_get_cache on a
|
|
throwaway event loop, reusing an async Redis client created on another loop and
|
|
corrupting its connection pool. The sync path must use the blocking client, never
|
|
the async one, and never create an event loop, even when called from a coroutine
|
|
(e.g. async_raise_no_deployment_exception -> get_min_cooldown).
|
|
"""
|
|
mock_redis = _redis_mock_for_sync_batch({"lit6729_key": "redis_value"})
|
|
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis)
|
|
|
|
_assert_sync_batch_used_blocking_client(dual_cache, mock_redis)
|
|
|
|
|
|
def test_dual_cache_batch_get_cache_uses_sync_redis_client_without_running_loop():
|
|
mock_redis = _redis_mock_for_sync_batch({"lit6729_key": "redis_value"})
|
|
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis)
|
|
|
|
_assert_sync_batch_used_blocking_client(dual_cache, mock_redis)
|
|
|
|
|
|
def test_dual_cache_batch_get_cache_only_reads_missing_keys_from_redis():
|
|
mock_redis = _redis_mock_for_sync_batch({"miss_key": "from_redis"})
|
|
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis)
|
|
dual_cache.in_memory_cache.set_cache("hit_key", "from_memory")
|
|
|
|
result = dual_cache.batch_get_cache(keys=["hit_key", "miss_key"])
|
|
|
|
assert result == ["from_memory", "from_redis"]
|
|
mock_redis.batch_get_cache.assert_called_once_with(key_list=["miss_key"], parent_otel_span=None)
|
|
|
|
|
|
def test_dual_cache_batch_get_cache_throttles_repeat_redis_reads():
|
|
mock_redis = _redis_mock_for_sync_batch({"absent_key": None})
|
|
dual_cache = DualCache(
|
|
in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10
|
|
)
|
|
|
|
first = dual_cache.batch_get_cache(keys=["absent_key"])
|
|
second = dual_cache.batch_get_cache(keys=["absent_key"])
|
|
|
|
assert first == [None]
|
|
assert second == [None]
|
|
mock_redis.batch_get_cache.assert_called_once()
|
|
|
|
|
|
def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error():
|
|
mock_redis = MagicMock(spec=RedisCache)
|
|
mock_redis.batch_get_cache.side_effect = RuntimeError("redis unavailable")
|
|
dual_cache = DualCache(
|
|
in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10
|
|
)
|
|
|
|
first_result = dual_cache.batch_get_cache(keys=["shared_a"])
|
|
second_result = dual_cache.batch_get_cache(keys=["shared_a"])
|
|
|
|
assert first_result is None
|
|
assert second_result is None
|
|
assert mock_redis.batch_get_cache.call_count == 2
|
|
assert "shared_a" not in dual_cache.last_redis_batch_access_time
|
|
|
|
|
|
def test_dual_cache_batch_get_cache_returns_memory_only_when_redis_read_is_throttled():
|
|
mock_redis = _redis_mock_for_sync_batch({"throttled_key": "redis_value"})
|
|
dual_cache = DualCache(
|
|
in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10
|
|
)
|
|
dual_cache.last_redis_batch_access_time["throttled_key"] = time.time()
|
|
|
|
result = dual_cache.batch_get_cache(keys=["throttled_key"])
|
|
|
|
assert result == [None]
|
|
mock_redis.batch_get_cache.assert_not_called()
|
|
|
|
|
|
def test_dual_cache_sync_batch_redis_backfill_injects_default_in_memory_ttl():
|
|
"""Sync batch_get_cache's Redis-to-memory backfill must honor
|
|
default_in_memory_ttl, same as the async path."""
|
|
in_memory_cache = InMemoryCache(default_ttl=600)
|
|
mock_redis = _redis_mock_for_sync_batch({"batch_backfill_key": "redis_value"})
|
|
dual_cache = DualCache(
|
|
in_memory_cache=in_memory_cache,
|
|
redis_cache=mock_redis,
|
|
default_in_memory_ttl=60,
|
|
)
|
|
|
|
before = time.time()
|
|
result = dual_cache.batch_get_cache(keys=["batch_backfill_key"])
|
|
after = time.time()
|
|
|
|
assert result == ["redis_value"]
|
|
expiry = in_memory_cache.ttl_dict["batch_backfill_key"]
|
|
assert expiry >= before + 60
|
|
assert expiry <= after + 60
|
|
|
|
|
|
def test_dual_cache_batch_get_cache_forwards_explicit_ttl_to_backfill():
|
|
"""An explicit ttl kwarg must reach the in-memory backfill flat, not nested
|
|
under a 'kwargs' key the way the old locals()-forwarding path sent it."""
|
|
in_memory_cache = InMemoryCache(default_ttl=600)
|
|
mock_redis = _redis_mock_for_sync_batch({"explicit_ttl_key": "redis_value"})
|
|
dual_cache = DualCache(in_memory_cache=in_memory_cache, redis_cache=mock_redis)
|
|
|
|
before = time.time()
|
|
result = dual_cache.batch_get_cache(keys=["explicit_ttl_key"], ttl=5)
|
|
after = time.time()
|
|
|
|
assert result == ["redis_value"]
|
|
expiry = in_memory_cache.ttl_dict["explicit_ttl_key"]
|
|
assert expiry >= before + 5
|
|
assert expiry <= after + 5
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_async_set_cache_injects_default_in_memory_ttl():
|
|
"""
|
|
Test that async_set_cache injects default_in_memory_ttl into kwargs
|
|
when no explicit ttl is provided, matching the sync set_cache behavior.
|
|
|
|
Regression test for: async_set_cache was missing the TTL injection that
|
|
sync set_cache has, causing InMemoryCache to use its own default_ttl (600s)
|
|
instead of DualCache's default_in_memory_ttl.
|
|
"""
|
|
in_memory_cache = InMemoryCache(default_ttl=600)
|
|
dual_cache = DualCache(
|
|
in_memory_cache=in_memory_cache,
|
|
default_in_memory_ttl=60,
|
|
)
|
|
|
|
before = time.time()
|
|
await dual_cache.async_set_cache(key="test_key", value="test_value")
|
|
after = time.time()
|
|
|
|
# The TTL stored should reflect default_in_memory_ttl (60s), not
|
|
# InMemoryCache's default_ttl (600s)
|
|
expiry = in_memory_cache.ttl_dict["test_key"]
|
|
assert expiry >= before + 60
|
|
assert expiry <= after + 60
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_redis_backfill_injects_default_in_memory_ttl():
|
|
"""
|
|
A Redis-hit backfill into the in-memory tier must honor
|
|
default_in_memory_ttl the same way the write paths do. Without it, the
|
|
backfilled entry falls to InMemoryCache's own default_ttl (600s), so a
|
|
replica that primed a management object (e.g. a virtual key's auth blob)
|
|
from Redis keeps serving it for 10 minutes after the object was updated
|
|
and invalidated, instead of re-reading within the configured TTL.
|
|
"""
|
|
in_memory_cache = InMemoryCache(default_ttl=600)
|
|
redis_cache = MagicMock()
|
|
redis_cache.async_get_cache = AsyncMock(return_value="redis_value")
|
|
dual_cache = DualCache(
|
|
in_memory_cache=in_memory_cache,
|
|
redis_cache=redis_cache,
|
|
default_in_memory_ttl=60,
|
|
)
|
|
|
|
before = time.time()
|
|
result = await dual_cache.async_get_cache(key="backfill_key")
|
|
after = time.time()
|
|
|
|
assert result == "redis_value"
|
|
expiry = in_memory_cache.ttl_dict["backfill_key"]
|
|
assert expiry >= before + 60
|
|
assert expiry <= after + 60
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_batch_redis_backfill_injects_default_in_memory_ttl():
|
|
"""async_batch_get_cache's Redis-to-memory backfill must honor
|
|
default_in_memory_ttl, same as the single-key path."""
|
|
in_memory_cache = InMemoryCache(default_ttl=600)
|
|
mock_redis = MagicMock(spec=RedisCache)
|
|
mock_redis.async_batch_get_cache = AsyncMock(
|
|
return_value={"batch_backfill_key": "redis_value"}
|
|
)
|
|
dual_cache = DualCache(
|
|
in_memory_cache=in_memory_cache,
|
|
redis_cache=mock_redis,
|
|
default_in_memory_ttl=60,
|
|
)
|
|
|
|
before = time.time()
|
|
result = await dual_cache.async_batch_get_cache(keys=["batch_backfill_key"])
|
|
after = time.time()
|
|
|
|
assert result == ["redis_value"]
|
|
expiry = in_memory_cache.ttl_dict["batch_backfill_key"]
|
|
assert expiry >= before + 60
|
|
assert expiry <= after + 60
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_async_set_cache_respects_explicit_ttl():
|
|
"""
|
|
Test that async_set_cache does NOT override an explicitly provided ttl.
|
|
"""
|
|
in_memory_cache = InMemoryCache(default_ttl=600)
|
|
dual_cache = DualCache(
|
|
in_memory_cache=in_memory_cache,
|
|
default_in_memory_ttl=60,
|
|
)
|
|
|
|
before = time.time()
|
|
await dual_cache.async_set_cache(key="test_key", value="test_value", ttl=30)
|
|
after = time.time()
|
|
|
|
# The explicit ttl=30 should be used, not default_in_memory_ttl (60)
|
|
expiry = in_memory_cache.ttl_dict["test_key"]
|
|
assert expiry >= before + 30
|
|
assert expiry <= after + 30
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_async_set_cache_pipeline_injects_default_in_memory_ttl():
|
|
"""
|
|
Test that async_set_cache_pipeline injects default_in_memory_ttl into kwargs
|
|
when no explicit ttl is provided.
|
|
"""
|
|
in_memory_cache = InMemoryCache(default_ttl=600)
|
|
dual_cache = DualCache(
|
|
in_memory_cache=in_memory_cache,
|
|
default_in_memory_ttl=60,
|
|
)
|
|
|
|
cache_list = [("key_a", "value_a"), ("key_b", "value_b")]
|
|
|
|
before = time.time()
|
|
await dual_cache.async_set_cache_pipeline(cache_list=cache_list)
|
|
after = time.time()
|
|
|
|
for key in ["key_a", "key_b"]:
|
|
expiry = in_memory_cache.ttl_dict[key]
|
|
assert expiry >= before + 60
|
|
assert expiry <= after + 60
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_sync_and_async_set_cache_use_same_ttl():
|
|
"""
|
|
Test that sync set_cache and async async_set_cache produce the same TTL
|
|
when no explicit ttl is provided, ensuring parity between the two paths.
|
|
"""
|
|
in_memory_sync = InMemoryCache(default_ttl=600)
|
|
dual_cache_sync = DualCache(
|
|
in_memory_cache=in_memory_sync,
|
|
default_in_memory_ttl=60,
|
|
)
|
|
|
|
in_memory_async = InMemoryCache(default_ttl=600)
|
|
dual_cache_async = DualCache(
|
|
in_memory_cache=in_memory_async,
|
|
default_in_memory_ttl=60,
|
|
)
|
|
|
|
dual_cache_sync.set_cache(key="test_key", value="test_value")
|
|
await dual_cache_async.async_set_cache(key="test_key", value="test_value")
|
|
|
|
sync_expiry = in_memory_sync.ttl_dict["test_key"]
|
|
async_expiry = in_memory_async.ttl_dict["test_key"]
|
|
|
|
# Both should use default_in_memory_ttl=60, so their expiry times
|
|
# should be within a small tolerance of each other
|
|
assert abs(sync_expiry - async_expiry) < 1.0
|
|
|
|
|
|
def test_circuit_breaker_opens_after_threshold():
|
|
"""Circuit opens after N consecutive Redis failures."""
|
|
from litellm.caching.redis_cache import RedisCircuitBreaker
|
|
|
|
cb = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
|
for _ in range(3):
|
|
cb.record_failure()
|
|
|
|
assert cb._state == "open"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_circuit_breaker_open_skips_redis():
|
|
"""When circuit is open, the guard decorator raises immediately without calling the method."""
|
|
from litellm.caching.redis_cache import (
|
|
RedisCircuitBreaker,
|
|
_redis_circuit_breaker_guard,
|
|
)
|
|
|
|
class FakeRedis:
|
|
def __init__(self):
|
|
self._circuit_breaker = RedisCircuitBreaker(
|
|
failure_threshold=3, recovery_timeout=60
|
|
)
|
|
self._circuit_breaker._state = "open"
|
|
self._circuit_breaker._opened_at = time.time()
|
|
self.call_count = 0
|
|
|
|
@_redis_circuit_breaker_guard
|
|
async def do_thing(self):
|
|
self.call_count += 1
|
|
return "result"
|
|
|
|
fr = FakeRedis()
|
|
with pytest.raises(Exception, match="circuit breaker is open"):
|
|
await fr.do_thing()
|
|
|
|
assert fr.call_count == 0 # method body never executed
|
|
|
|
|
|
def test_circuit_breaker_closes_on_recovery():
|
|
"""After recovery_timeout expires, probe is allowed and success closes the circuit."""
|
|
from litellm.caching.redis_cache import RedisCircuitBreaker
|
|
|
|
cb = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
|
cb._state = "open"
|
|
cb._opened_at = time.time() - 9999 # recovery timeout long expired
|
|
|
|
# is_open() should return False to allow a probe through, and transition to HALF_OPEN
|
|
assert cb.is_open() is False
|
|
assert cb._state == "half_open"
|
|
|
|
# Successful probe closes the circuit
|
|
cb.record_success()
|
|
assert cb._state == "closed"
|
|
|
|
|
|
def test_circuit_breaker_half_open_concurrent_calls_are_fast_failed():
|
|
"""
|
|
Regression test: only ONE probe gets through when the circuit transitions
|
|
OPEN → HALF_OPEN. All concurrent callers that check is_open() while the
|
|
state is already HALF_OPEN must be fast-failed (return True), not allowed
|
|
through as additional probes.
|
|
"""
|
|
from litellm.caching.redis_cache import RedisCircuitBreaker
|
|
|
|
cb = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
|
cb._state = "open"
|
|
cb._opened_at = time.time() - 9999 # recovery timeout long expired
|
|
|
|
# First caller: OPEN + expired → transitions to HALF_OPEN, returns False (probe)
|
|
assert cb.is_open() is False
|
|
assert cb._state == "half_open"
|
|
|
|
# All subsequent concurrent callers: HALF_OPEN → fast-fail (return True)
|
|
for _ in range(10):
|
|
assert (
|
|
cb.is_open() is True
|
|
), "concurrent callers should be fast-failed in HALF_OPEN"
|
|
|
|
|
|
def test_circuit_breaker_disabled_never_opens():
|
|
"""When disabled, failures never open the circuit and is_open() stays False."""
|
|
from litellm.caching.redis_cache import RedisCircuitBreaker
|
|
|
|
cb = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, enabled=False)
|
|
|
|
for _ in range(100):
|
|
cb.record_failure()
|
|
|
|
assert cb._state == "closed"
|
|
assert cb.is_open() is False
|
|
|
|
|
|
def test_circuit_breaker_disabled_record_success_leaves_state_untouched():
|
|
"""
|
|
A disabled breaker must not mutate state in any state-machine method. Force
|
|
a non-default (OPEN) state and assert record_success() returns without
|
|
resetting it — the same enabled-guard contract as is_open/record_failure.
|
|
"""
|
|
from litellm.caching.redis_cache import RedisCircuitBreaker
|
|
|
|
cb = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, enabled=False)
|
|
cb._state = "open"
|
|
cb._failure_count = 3
|
|
|
|
cb.record_success()
|
|
|
|
assert cb._state == "open"
|
|
assert cb._failure_count == 3
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_circuit_breaker_disabled_guard_always_calls_method():
|
|
"""A disabled breaker lets every guarded call through, even after failures."""
|
|
from litellm.caching.redis_cache import (
|
|
RedisCircuitBreaker,
|
|
_redis_circuit_breaker_guard,
|
|
)
|
|
|
|
class FakeRedis:
|
|
def __init__(self):
|
|
self._circuit_breaker = RedisCircuitBreaker(
|
|
failure_threshold=1, recovery_timeout=60, enabled=False
|
|
)
|
|
self.call_count = 0
|
|
|
|
@_redis_circuit_breaker_guard
|
|
async def boom(self):
|
|
self.call_count += 1
|
|
raise RuntimeError("redis down")
|
|
|
|
fr = FakeRedis()
|
|
for _ in range(5):
|
|
with pytest.raises(RuntimeError, match="redis down"):
|
|
await fr.boom()
|
|
|
|
# Every call reached the method body; the breaker never short-circuited.
|
|
assert fr.call_count == 5
|
|
assert fr._circuit_breaker.is_open() is False
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_increment_cache_returns_none_when_no_in_memory_cache_and_redis_fails():
|
|
"""
|
|
Regression test: when in_memory_cache is None and Redis fails, async_increment_cache
|
|
must return None — not the raw increment delta — to avoid silently miscalculating
|
|
rate-limit counters.
|
|
"""
|
|
dc = DualCache()
|
|
dc.in_memory_cache = None # type: ignore[assignment] # constructor always creates InMemoryCache, so null it manually
|
|
dc.redis_cache = MagicMock()
|
|
dc.redis_cache.async_increment = AsyncMock(side_effect=Exception("redis down"))
|
|
|
|
result = await dc.async_increment_cache("rpm:model:14-05", 1.0, ttl=60)
|
|
|
|
assert result is None, (
|
|
f"Expected None when in_memory_cache is absent and Redis fails, got {result!r}. "
|
|
"Returning the delta (1.0) would silently miscalculate rate-limit counters."
|
|
)
|
|
|
|
|
|
def test_dual_cache_late_attach_redis_wires_writes_and_ttl_sync():
|
|
"""
|
|
Typical lazy startup (sync): DualCache runs with in-memory only, then Redis
|
|
becomes available and is attached. New writes must reach Redis; keys written
|
|
before attach are not backfilled. Optional default_redis_ttl is applied on attach.
|
|
"""
|
|
in_memory = InMemoryCache()
|
|
dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None)
|
|
|
|
mock_redis = MagicMock()
|
|
mock_redis.set_cache = MagicMock()
|
|
mock_redis.async_set_cache = AsyncMock()
|
|
|
|
key_before = f"before_attach_{uuid.uuid4()}"
|
|
val_before = {"phase": "memory_only"}
|
|
dual_cache.set_cache(key_before, val_before)
|
|
|
|
assert in_memory.get_cache(key_before) == val_before
|
|
|
|
dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0)
|
|
assert dual_cache.redis_cache is mock_redis
|
|
assert dual_cache.default_redis_ttl == 99.0
|
|
|
|
mock_redis.set_cache.assert_not_called()
|
|
|
|
key_after = f"after_attach_{uuid.uuid4()}"
|
|
val_after = {"phase": "memory_and_redis"}
|
|
dual_cache.set_cache(key_after, val_after)
|
|
mock_redis.set_cache.assert_called_once()
|
|
assert mock_redis.set_cache.call_args[0][:2] == (key_after, val_after)
|
|
|
|
assert in_memory.get_cache(key_after) == val_after
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dual_cache_late_attach_redis_wires_writes_and_ttl_async():
|
|
"""
|
|
Typical lazy startup (async): DualCache runs with in-memory only, then Redis
|
|
becomes available and is attached. New writes must reach Redis; keys written
|
|
before attach are not backfilled. Optional default_redis_ttl is applied on attach.
|
|
"""
|
|
in_memory = InMemoryCache()
|
|
dual_cache = DualCache(in_memory_cache=in_memory, redis_cache=None)
|
|
|
|
mock_redis = MagicMock()
|
|
mock_redis.set_cache = MagicMock()
|
|
mock_redis.async_set_cache = AsyncMock()
|
|
|
|
key_before = f"before_attach_{uuid.uuid4()}"
|
|
val_before = {"phase": "memory_only"}
|
|
await dual_cache.async_set_cache(key_before, val_before)
|
|
|
|
assert in_memory.get_cache(key_before) == val_before
|
|
|
|
dual_cache.attach_redis_cache(mock_redis, default_redis_ttl=99.0)
|
|
assert dual_cache.redis_cache is mock_redis
|
|
assert dual_cache.default_redis_ttl == 99.0
|
|
|
|
mock_redis.async_set_cache.assert_not_called()
|
|
|
|
key_after = f"after_attach_{uuid.uuid4()}"
|
|
val_after = {"phase": "memory_and_redis"}
|
|
await dual_cache.async_set_cache(key_after, val_after)
|
|
mock_redis.async_set_cache.assert_called_once()
|
|
assert mock_redis.async_set_cache.call_args[0][:2] == (key_after, val_after)
|
|
|
|
assert in_memory.get_cache(key_after) == val_after
|
|
|
|
|
|
class _OpenBreakerRedis:
|
|
def __init__(self) -> None:
|
|
from litellm.caching.redis_cache import RedisCircuitBreaker
|
|
|
|
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
|
|
for _ in range(3):
|
|
self._circuit_breaker.record_failure()
|
|
|
|
@_redis_circuit_breaker_guard
|
|
async def async_get_cache(self, key, **kwargs):
|
|
raise AssertionError("never reached")
|
|
|
|
@_redis_circuit_breaker_guard
|
|
async def async_batch_get_cache(self, key_list, **kwargs):
|
|
raise AssertionError("never reached")
|
|
|
|
@_redis_circuit_breaker_guard
|
|
async def async_set_cache(self, key, value, **kwargs):
|
|
raise AssertionError("never reached")
|
|
|
|
@_redis_circuit_breaker_guard
|
|
async def async_set_cache_pipeline(self, cache_list, **kwargs):
|
|
raise AssertionError("never reached")
|
|
|
|
@_redis_circuit_breaker_guard
|
|
async def async_increment_pipeline(self, increment_list, **kwargs):
|
|
raise AssertionError("never reached")
|
|
|
|
@_redis_circuit_breaker_guard
|
|
async def async_increment(self, key, value, **kwargs):
|
|
raise AssertionError("never reached")
|
|
|
|
@_redis_circuit_breaker_guard_sync
|
|
def get_cache(self, key, **kwargs):
|
|
raise AssertionError("never reached")
|
|
|
|
@_redis_circuit_breaker_guard_sync
|
|
def batch_get_cache(self, key_list, **kwargs):
|
|
raise AssertionError("never reached")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"call",
|
|
[
|
|
lambda cache: cache.async_get_cache("k"),
|
|
lambda cache: cache.async_batch_get_cache(["k1", "k2"]),
|
|
lambda cache: cache.async_set_cache("k", "v"),
|
|
lambda cache: cache.async_set_cache_pipeline([("k", "v")]),
|
|
lambda cache: cache.async_increment_cache_pipeline(
|
|
increment_list=[RedisPipelineIncrementOperation(key="k", increment_value=1.0, ttl=60)]
|
|
),
|
|
lambda cache: cache.async_increment_cache("k", 1.0),
|
|
],
|
|
ids=["get", "batch_get", "set", "set_pipeline", "increment_pipeline", "increment"],
|
|
)
|
|
async def test_an_open_circuit_breaker_is_not_an_error_per_request(caplog, call):
|
|
cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
|
caplog.clear()
|
|
|
|
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
|
await call(cache)
|
|
|
|
assert [record.levelno for record in caplog.records if record.levelno >= logging.WARNING] == []
|
|
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"call",
|
|
[lambda cache: cache.get_cache("k"), lambda cache: cache.batch_get_cache(["k1", "k2"])],
|
|
ids=["get", "batch_get"],
|
|
)
|
|
def test_an_open_circuit_breaker_is_not_an_error_per_sync_request(caplog, call):
|
|
cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=_OpenBreakerRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
|
caplog.clear()
|
|
|
|
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
|
call(cache)
|
|
|
|
assert [record.levelno for record in caplog.records if record.levelno >= logging.WARNING] == []
|
|
assert any("circuit breaker is open" in record.getMessage() for record in caplog.records)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_a_real_redis_failure_still_logs_an_error(caplog):
|
|
class _BrokenRedis:
|
|
async def async_get_cache(self, key, **kwargs):
|
|
raise ConnectionError("redis is down")
|
|
|
|
cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=_BrokenRedis()) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
|
|
|
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
|
assert await cache.async_get_cache("k") is None
|
|
|
|
errors = [record for record in caplog.records if record.levelno == logging.ERROR]
|
|
assert [record.getMessage() for record in errors] == ["LiteLLM Cache: exception in async_get_cache: redis is down"]
|
|
assert errors[0].exc_info is not None
|
|
|
|
|
|
def _dual_cache_with_open_breaker_and_a_memory_hit() -> DualCache:
|
|
in_memory = InMemoryCache()
|
|
in_memory.set_cache("k1", "v1")
|
|
return DualCache(in_memory_cache=in_memory, redis_cache=_OpenBreakerRedis(), default_redis_batch_cache_expiry=10) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
|
|
|
|
|
def test_open_breaker_keeps_sync_batch_read_memory_hits_and_releases_reservations():
|
|
"""A refused Redis batch read must still answer with the in-memory hits and hold no reservation.
|
|
|
|
The refusal was logged and turned into a bare None, so a caller lost its in-memory hits
|
|
for as long as the breaker stayed open, and the reserved keys stayed throttled until
|
|
the batch expiry passed even though nothing was ever read for them.
|
|
"""
|
|
cache = _dual_cache_with_open_breaker_and_a_memory_hit()
|
|
|
|
assert list(cache.batch_get_cache(["k1", "k2"])) == ["v1", None]
|
|
assert "k2" not in cache.last_redis_batch_access_time
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_open_breaker_keeps_async_batch_read_memory_hits_and_releases_reservations():
|
|
cache = _dual_cache_with_open_breaker_and_a_memory_hit()
|
|
|
|
assert list(await cache.async_batch_get_cache(["k1", "k2"])) == ["v1", None]
|
|
assert "k2" not in cache.last_redis_batch_access_time
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_redis_timeouts_falling_back_to_memory_log_once_per_interval(caplog, monkeypatch):
|
|
"""The first fallback WARNING of a timeout streak logs, the rest stay at DEBUG until the summary."""
|
|
from redis.exceptions import TimeoutError as RedisTimeoutError
|
|
|
|
from litellm.caching import redis_cache as redis_cache_module
|
|
from litellm.caching.redis_cache import _RedisTimeoutLogThrottle
|
|
|
|
clock = MagicMock(return_value=1_000.0)
|
|
monkeypatch.setattr(
|
|
redis_cache_module, "_redis_timeout_log_throttle", _RedisTimeoutLogThrottle(interval=5.0, clock=clock)
|
|
)
|
|
|
|
class _TimingOutRedis:
|
|
async def async_increment_pipeline(self, increment_list, **kwargs):
|
|
raise RedisTimeoutError("Timeout reading from 127.0.0.1:6379")
|
|
|
|
async def async_increment(self, key, value, **kwargs):
|
|
raise RedisTimeoutError("Timeout reading from 127.0.0.1:6379")
|
|
|
|
cache = DualCache(
|
|
in_memory_cache=InMemoryCache(),
|
|
redis_cache=_TimingOutRedis(), # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
|
)
|
|
increments = [RedisPipelineIncrementOperation(key="k", increment_value=1.0, ttl=60)]
|
|
|
|
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
|
for _ in range(100):
|
|
await cache.async_increment_cache_pipeline(increment_list=increments)
|
|
await cache.async_increment_cache("k", 1.0)
|
|
|
|
visible = [r for r in caplog.records if r.levelno >= logging.WARNING]
|
|
assert [(r.levelno, r.getMessage()) for r in visible] == [
|
|
(
|
|
logging.WARNING,
|
|
"Redis async_increment_cache_pipeline failed, falling back to in-memory result:"
|
|
" Timeout reading from 127.0.0.1:6379",
|
|
)
|
|
]
|
|
assert visible[0].filename == "dual_cache.py"
|
|
assert sum("Timeout reading from" in r.getMessage() for r in caplog.records) == 200
|
|
|
|
caplog.clear()
|
|
clock.return_value += 5.0
|
|
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
|
await cache.async_increment_cache("k", 1.0)
|
|
assert [(r.levelno, r.getMessage()) for r in caplog.records] == [
|
|
(
|
|
logging.WARNING,
|
|
"Redis async_increment_cache failed, falling back to in-memory result: Timeout reading from 127.0.0.1:6379"
|
|
" (199 more Redis timeouts since the previous Redis timeout line were logged at DEBUG)",
|
|
)
|
|
]
|