mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor(cache): keep native foundation isolated
This commit is contained in:
parent
595711829a
commit
1b6b704ddd
9 changed files with 42 additions and 160 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -22,8 +22,5 @@
|
|||
],
|
||||
"secret_manager": [
|
||||
"readable"
|
||||
],
|
||||
"cache_settings": [
|
||||
"default_redis_ttl"
|
||||
]
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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: ...
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue