refactor(cache): keep native foundation isolated

This commit is contained in:
Yujong Lee 2026-09-21 11:15:44 -07:00
parent 595711829a
commit 1b6b704ddd
9 changed files with 42 additions and 160 deletions

View file

@ -42,7 +42,7 @@ The resolver reads the namespace's `cache` attribute each time it resolves. A ca
Python callbacks use the built-in `Cache` API, so a `Cache` subclass works unchanged. A batch lookup takes one original kwargs mapping per request and returns the list of `get_cache` or gathered `async_get_cache` results, while native bindings return `{values, missing_indices}`. A batch store hands the caller's original result to `async_add_cache_pipeline`. `ping` calls `ping`, and a flush goes to the facade's backend
The private facade test harness checks object identity, method overrides, effective TTL, Redis namespace, memory capacity, and later configuration changes before selecting native execution. Its snapshot includes Redis connection settings, so a later `redis_kwargs` change, including an SSL option, selects Python callback execution. Redis defaults come from the Python settings snapshot, including `litellm.default_redis_ttl`, and buffered async writes honor `redis_flush_size`. Public activation must construct the shared native service from the initial Python Redis settings. A buffered entry keeps the time it was produced, and a failed flush drops its batch instead of growing the buffer during an outage. The harness does not migrate entries or replace Python methods. Until activation configures one shared service, the Python facade and native test service can hold separate data. Existing public cache constructors remain on Python
The private facade test harness checks object identity, method overrides, effective TTL, Redis namespace, memory capacity, and later configuration changes before selecting native execution. Its snapshot includes Redis connection settings, so a later `redis_kwargs` change, including an SSL option, selects Python callback execution. Buffered async writes honor `redis_flush_size`. Public activation must construct the shared native service from the initial Python Redis settings, including `litellm.default_redis_ttl` and SSL options. A buffered entry keeps the time it was produced, and a failed flush drops its batch instead of growing the buffer during an outage. The harness does not migrate entries or replace Python methods. Until activation configures one shared service, the Python facade and native test service can hold separate data. Existing public cache constructors remain on Python
Native cache handles must be recreated after fork. The bridge releases the GIL around native operations, and Redis runs blocking connection operations off the async executor. Native errors propagate to the host, which owns the existing fail-open and logging policy

View file

@ -22,8 +22,5 @@
],
"secret_manager": [
"readable"
],
"cache_settings": [
"default_redis_ttl"
]
}

View file

@ -1,26 +1,7 @@
use std::time::Duration;
use litellm_host_python::release_gil;
use pyo3::{PyTraverseError, PyVisit, exceptions::PyRuntimeError, prelude::*};
use super::{cache_error, facade::FacadeGuard, native::NativeResponseCache, request::duration};
use crate::python_settings::PythonSettings;
const PYTHON_REDIS_DEFAULT_TTL: Duration = Duration::from_secs(60);
#[derive(FromPyObject)]
struct PythonCacheSettings {
default_redis_ttl: Option<f64>,
}
fn redis_default_ttl(py: Python<'_>) -> PyResult<Duration> {
let settings: PythonCacheSettings = PythonSettings::Cache.read(py)?.extract()?;
settings
.default_redis_ttl
.map(duration)
.transpose()
.map(|ttl| ttl.unwrap_or(PYTHON_REDIS_DEFAULT_TTL))
}
#[pyclass(frozen, name = "_CacheTestHandle")]
pub(crate) struct CacheTestHandle {
@ -53,17 +34,14 @@ impl CacheTestHandle {
}
#[staticmethod]
#[pyo3(signature = (url, *, ttl_seconds=None, namespace=None))]
#[pyo3(signature = (url, *, ttl_seconds=60.0, namespace=None))]
fn redis(
py: Python<'_>,
url: String,
ttl_seconds: Option<f64>,
ttl_seconds: f64,
namespace: Option<String>,
) -> PyResult<Self> {
let ttl = Some(match ttl_seconds {
Some(seconds) => duration(seconds)?,
None => redis_default_ttl(py)?,
});
let ttl = Some(duration(ttl_seconds)?);
let service = release_gil(py, move || NativeResponseCache::redis(&url, ttl, namespace))
.map_err(cache_error)?;
Ok(Self {

View file

@ -8,17 +8,15 @@ pub(crate) enum PythonSettings {
UrlPolicy,
ProviderDefaults,
SecretManager,
Cache,
}
impl PythonSettings {
#[cfg(test)]
pub(crate) const ALL: [Self; 5] = [
pub(crate) const ALL: [Self; 4] = [
Self::Http,
Self::UrlPolicy,
Self::ProviderDefaults,
Self::SecretManager,
Self::Cache,
];
pub(crate) fn name(self) -> &'static str {
@ -27,7 +25,6 @@ impl PythonSettings {
Self::UrlPolicy => "url_policy",
Self::ProviderDefaults => "provider_defaults",
Self::SecretManager => "secret_manager",
Self::Cache => "cache_settings",
}
}

View file

@ -12,7 +12,7 @@ import logging
import time
from collections.abc import Sequence
from threading import Lock
from typing import TYPE_CHECKING, Any, Final, TypeVar
from typing import TYPE_CHECKING, Any, Final
if TYPE_CHECKING:
from litellm.types.caching import RedisPipelineIncrementOperation
@ -34,20 +34,14 @@ else:
from collections import OrderedDict
_KeyT = TypeVar("_KeyT")
_ValueT = TypeVar("_ValueT")
class LimitedSizeOrderedDict(OrderedDict[_KeyT, _ValueT]):
def __init__(self, *, max_size: int = 100) -> None:
super().__init__()
class LimitedSizeOrderedDict(OrderedDict):
def __init__(self, *args, max_size=100, **kwargs):
super().__init__(*args, **kwargs)
self.max_size = max_size
def __setitem__(self, key: _KeyT, value: _ValueT) -> None:
if key in self:
super().__setitem__(key, value)
self.move_to_end(key)
return
def __setitem__(self, key, value):
# If inserting a new key exceeds max size, remove the oldest item
if len(self) >= self.max_size:
self.popitem(last=False)
super().__setitem__(key, value)
@ -74,9 +68,7 @@ class DualCache(BaseCache):
self.in_memory_cache = in_memory_cache or InMemoryCache()
# If redis_cache is not provided, use the default RedisCache
self.redis_cache = redis_cache
self.last_redis_batch_access_time: LimitedSizeOrderedDict[str, float] = LimitedSizeOrderedDict(
max_size=default_max_redis_batch_cache_size
)
self.last_redis_batch_access_time = LimitedSizeOrderedDict(max_size=default_max_redis_batch_cache_size)
self._last_redis_batch_access_time_lock = Lock()
self.redis_batch_cache_expiry = (
default_redis_batch_cache_expiry or litellm.default_redis_batch_cache_expiry or 10

View file

@ -1,5 +1,5 @@
from asyncio import Future
from collections.abc import AsyncIterator, Awaitable, Coroutine, Iterator, Mapping, Sequence
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence
from typing import Never, final
from litellm.llms.base_llm.ocr.transformation import OCRResponse
@ -93,72 +93,6 @@ class ResponsesWebSocketConnection:
def recv_text(self) -> Future[str | None]: ...
def close(self) -> Future[None]: ...
@final
class _CacheTestHandle:
def __new__(cls, _uninstantiable: Never, /) -> Never: ...
@staticmethod
def memory(
*, capacity: int = 200, ttl_seconds: float = 600.0, max_entry_bytes: int = 1048576
) -> _CacheTestHandle: ...
@staticmethod
def redis(url: str, *, ttl_seconds: float | None = None, namespace: str | None = None) -> _CacheTestHandle: ...
@property
def backend(self) -> str: ...
def _bind_facade(self, facade: object) -> None: ...
@final
class _CacheTestResolver:
def __new__(cls, namespace: object) -> _CacheTestResolver: ...
def resolve(self) -> _CacheTestBinding: ...
@final
class _CacheTestBinding:
def __new__(cls, _uninstantiable: Never, /) -> Never: ...
@property
def kind(self) -> str: ...
def lookup(
self, request: Mapping[str, object] | None, *, callback_kwargs: dict[str, object] | None = None
) -> object: ...
def store(
self,
request: Mapping[str, object] | None,
response: object,
*,
callback_kwargs: dict[str, object] | None = None,
) -> None: ...
def lookup_batch(
self,
requests: Sequence[Mapping[str, object]],
*,
callback_kwargs: Sequence[dict[str, object]] | None = None,
) -> object: ...
def async_lookup(
self, request: Mapping[str, object] | None, *, callback_kwargs: dict[str, object] | None = None
) -> Awaitable[object]: ...
def async_store(
self,
request: Mapping[str, object] | None,
response: object,
*,
callback_kwargs: dict[str, object] | None = None,
) -> Awaitable[object]: ...
def async_lookup_batch(
self,
requests: Sequence[Mapping[str, object]],
*,
callback_kwargs: Sequence[dict[str, object]] | None = None,
) -> Awaitable[object]: ...
def async_store_batch(
self,
requests: Sequence[Mapping[str, object]],
responses: Sequence[object],
*,
callback_result: object = None,
callback_kwargs: dict[str, object] | None = None,
) -> Awaitable[object]: ...
def async_flush(self) -> Awaitable[None]: ...
def ping(self) -> Awaitable[object]: ...
@final
class TokenCounter:
def __new__(cls, tokenizer_json: str) -> TokenCounter: ...

View file

@ -36,11 +36,6 @@ class SecretManager:
readable: bool
@dataclass(frozen=True, slots=True)
class CacheSettings:
default_redis_ttl: float | None
def warn(message: str) -> None:
from litellm._logging import verbose_logger
@ -55,12 +50,6 @@ def secret_manager() -> SecretManager:
return SecretManager(readable=_should_read_secret_from_secret_manager())
def cache_settings() -> CacheSettings:
import litellm
return CacheSettings(default_redis_ttl=litellm.default_redis_ttl)
def provider_defaults() -> ProviderDefaults:
import litellm

View file

@ -7,7 +7,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE
from litellm.caching.dual_cache import DualCache, LimitedSizeOrderedDict
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
@ -15,7 +15,9 @@ 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)
dual_cache = DualCache(
redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10
)
keys = ["shared_a", "shared_b"]
start_gate = asyncio.Event()
@ -42,7 +44,9 @@ async def test_dual_cache_async_batch_get_cache_coalesces_concurrent_redis_reads
@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)
dual_cache = DualCache(
redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10
)
keys = ["shared_a", "shared_b"]
with patch.object(
@ -112,7 +116,9 @@ def test_dual_cache_batch_get_cache_only_reads_missing_keys_from_redis():
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)
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"])
@ -125,7 +131,9 @@ def test_dual_cache_batch_get_cache_throttles_repeat_redis_reads():
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)
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"])
@ -138,7 +146,9 @@ def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error():
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 = 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"])
@ -247,7 +257,9 @@ async def test_dual_cache_batch_redis_backfill_injects_default_in_memory_ttl():
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"})
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,
@ -359,7 +371,9 @@ async def test_circuit_breaker_open_skips_redis():
class FakeRedis:
def __init__(self):
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
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
@ -412,7 +426,9 @@ def test_circuit_breaker_half_open_concurrent_calls_are_fast_failed():
# 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"
assert (
cb.is_open() is True
), "concurrent callers should be fast-failed in HALF_OPEN"
def test_circuit_breaker_disabled_never_opens():
@ -456,7 +472,9 @@ async def test_circuit_breaker_disabled_guard_always_calls_method():
class FakeRedis:
def __init__(self):
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=1, recovery_timeout=60, enabled=False)
self._circuit_breaker = RedisCircuitBreaker(
failure_threshold=1, recovery_timeout=60, enabled=False
)
self.call_count = 0
@_redis_circuit_breaker_guard
@ -773,14 +791,3 @@ async def test_async_delete_cache_keys_on_empty_list_touches_no_backend():
await dual_cache.async_delete_cache_keys([])
redis_cache.delete_cache_keys.assert_not_awaited()
def test_limited_ordered_dict_refreshes_recency_without_evicting_another_key():
tracker = LimitedSizeOrderedDict(max_size=2)
tracker["hot"] = 1
tracker["cold"] = 2
tracker["hot"] = 3
tracker["new"] = 4
assert list(tracker.items()) == [("hot", 3), ("new", 4)]

View file

@ -361,18 +361,6 @@ def test_facade_registration_rejects_mismatched_capacity() -> None:
_native._CacheTestHandle.memory(capacity=7)._bind_facade(facade)
async def test_redis_handle_reads_the_python_default_ttl(redis_url: str) -> None:
client: Final = redis.Redis.from_url(redis_url)
with rebound(litellm, "default_redis_ttl", 7):
binding: Final = _native._CacheTestResolver(
SimpleNamespace(cache=_native._CacheTestHandle.redis(redis_url))
).resolve()
await binding.async_store(request("native-default"), {"value": 1})
assert 0 < client.ttl("native-default") <= 7
client.close()
async def test_redis_facade_buffers_native_async_writes(redis_url: str) -> None:
parsed: Final = urlparse(redis_url)
with rebound(litellm, "default_redis_ttl", 60):
@ -386,7 +374,7 @@ async def test_redis_facade_buffers_native_async_writes(redis_url: str) -> None:
_native._CacheTestHandle.redis(redis_url, ttl_seconds=61)._bind_facade(facade)
with pytest.raises(TypeError, match="namespaces must match"):
_native._CacheTestHandle.redis(redis_url, namespace="other")._bind_facade(facade)
_native._CacheTestHandle.redis(redis_url)._bind_facade(facade)
_native._CacheTestHandle.redis(redis_url, ttl_seconds=60)._bind_facade(facade)
binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
client: Final = redis.Redis.from_url(redis_url)