mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-25 01:02:15 +00:00
fix(cache): preserve batch callback contracts
This commit is contained in:
parent
9783b7a377
commit
da402b8aee
4 changed files with 102 additions and 67 deletions
|
|
@ -10,7 +10,7 @@ use pyo3::{
|
|||
PyTraverseError, PyVisit,
|
||||
exceptions::{PyRuntimeError, PyTypeError, PyValueError},
|
||||
prelude::*,
|
||||
types::PyDict,
|
||||
types::{PyDict, PyList},
|
||||
};
|
||||
use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
|
|
@ -309,7 +309,7 @@ impl ResolvedCache {
|
|||
.bind(py)
|
||||
.call_method(
|
||||
"batch_get_cache",
|
||||
(),
|
||||
(callback_keys(py, requests)?,),
|
||||
Some(self::callback_kwargs(callback_kwargs)?),
|
||||
)
|
||||
.map(Bound::unbind),
|
||||
|
|
@ -383,7 +383,7 @@ impl ResolvedCache {
|
|||
}
|
||||
CacheBinding::PythonCallback(object) => object.bind(py).call_method(
|
||||
"async_batch_get_cache",
|
||||
(),
|
||||
(callback_keys(py, requests)?,),
|
||||
Some(self::callback_kwargs(callback_kwargs)?),
|
||||
),
|
||||
}
|
||||
|
|
@ -416,11 +416,24 @@ impl ResolvedCache {
|
|||
cache_error,
|
||||
)
|
||||
}
|
||||
CacheBinding::PythonCallback(object) => object.bind(py).call_method(
|
||||
"async_set_cache_pipeline",
|
||||
(responses,),
|
||||
Some(self::callback_kwargs(callback_kwargs)?),
|
||||
),
|
||||
CacheBinding::PythonCallback(object) => {
|
||||
let keys = callback_keys(py, requests)?;
|
||||
let responses = responses.try_iter()?.collect::<PyResult<Vec<_>>>()?;
|
||||
if keys.len() != responses.len() {
|
||||
return Err(PyValueError::new_err(
|
||||
"batch cache requests and responses must have equal lengths",
|
||||
));
|
||||
}
|
||||
let cache_list = PyList::empty(py);
|
||||
for (key, response) in keys.iter().zip(responses) {
|
||||
cache_list.append((key, response))?;
|
||||
}
|
||||
object.bind(py).call_method(
|
||||
"async_set_cache_pipeline",
|
||||
(cache_list,),
|
||||
Some(self::callback_kwargs(callback_kwargs)?),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -471,6 +484,18 @@ fn callback_kwargs<'a, 'py>(
|
|||
})
|
||||
}
|
||||
|
||||
fn callback_keys<'py>(
|
||||
py: Python<'py>,
|
||||
requests: &Bound<'py, PyAny>,
|
||||
) -> PyResult<Bound<'py, PyList>> {
|
||||
PyList::new(
|
||||
py,
|
||||
self::requests(requests)?
|
||||
.into_iter()
|
||||
.map(|request| litellm_cache_response::cache_key(&request.key)),
|
||||
)
|
||||
}
|
||||
|
||||
fn ready_none(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
|
||||
ready_value(py, &())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -139,7 +139,7 @@ class DualCache(BaseCache):
|
|||
except Exception as e:
|
||||
print_verbose(e)
|
||||
|
||||
def increment_cache(self, key, value: int, local_only: bool = False, **kwargs) -> float:
|
||||
def increment_cache(self, key, value: int, local_only: bool = False, **kwargs) -> int:
|
||||
"""
|
||||
Key - the key in cache
|
||||
|
||||
|
|
@ -148,15 +148,14 @@ class DualCache(BaseCache):
|
|||
Returns - int - the incremented value
|
||||
"""
|
||||
try:
|
||||
if self.redis_cache is not None and local_only is False:
|
||||
result: Final = self.redis_cache.increment_cache(key, value, **kwargs)
|
||||
if self.in_memory_cache is not None:
|
||||
self.in_memory_cache.set_cache(key, result, **kwargs)
|
||||
return result
|
||||
|
||||
result: int = value
|
||||
if self.in_memory_cache is not None:
|
||||
return self.in_memory_cache.increment_cache(key, value, **kwargs)
|
||||
return value
|
||||
result = self.in_memory_cache.increment_cache(key, value, **kwargs)
|
||||
|
||||
if self.redis_cache is not None and local_only is False:
|
||||
result = self.redis_cache.increment_cache(key, value, **kwargs)
|
||||
|
||||
return result
|
||||
except Exception as e:
|
||||
verbose_logger.error("LiteLLM Cache: Excepton async add_cache: %s", e)
|
||||
raise e
|
||||
|
|
@ -430,30 +429,29 @@ class DualCache(BaseCache):
|
|||
Returns - the incremented value, or None if no cache backend is
|
||||
available (in_memory_cache is None and Redis failed/is absent).
|
||||
"""
|
||||
result: float | None = None
|
||||
try:
|
||||
if self.in_memory_cache is not None:
|
||||
result = await self.in_memory_cache.async_increment(key, value, **kwargs)
|
||||
|
||||
if self.redis_cache is not None and local_only is False:
|
||||
result: Final = await self.redis_cache.async_increment(
|
||||
result = await self.redis_cache.async_increment(
|
||||
key,
|
||||
value,
|
||||
parent_otel_span=parent_otel_span,
|
||||
ttl=kwargs.get("ttl", None),
|
||||
refresh_ttl=refresh_ttl,
|
||||
)
|
||||
if self.in_memory_cache is not None:
|
||||
await self.in_memory_cache.async_set_cache(key, result, **kwargs)
|
||||
return result
|
||||
|
||||
if self.in_memory_cache is not None:
|
||||
return await self.in_memory_cache.async_increment(key, value, **kwargs)
|
||||
return None
|
||||
return result
|
||||
except Exception as e:
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.WARNING,
|
||||
"Redis async_increment_cache failed; local counter unchanged",
|
||||
"Redis async_increment_cache failed, falling back to in-memory result",
|
||||
e,
|
||||
)
|
||||
return None
|
||||
return result
|
||||
|
||||
async def async_increment_cache_pipeline(
|
||||
self,
|
||||
|
|
@ -462,32 +460,29 @@ class DualCache(BaseCache):
|
|||
parent_otel_span: Span | None = None,
|
||||
**kwargs,
|
||||
) -> list[float] | None:
|
||||
result: list[float] | None = None
|
||||
try:
|
||||
if self.redis_cache is not None and local_only is False:
|
||||
result: Final = await self.redis_cache.async_increment_pipeline(
|
||||
increment_list=increment_list,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
if result is not None and self.in_memory_cache is not None:
|
||||
await self.in_memory_cache.async_set_cache_pipeline(
|
||||
cache_list=tuple((increment["key"], value) for increment, value in zip(increment_list, result))
|
||||
)
|
||||
return result
|
||||
|
||||
if self.in_memory_cache is not None:
|
||||
return await self.in_memory_cache.async_increment_pipeline(
|
||||
result = await self.in_memory_cache.async_increment_pipeline(
|
||||
increment_list=increment_list,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
return None
|
||||
|
||||
if self.redis_cache is not None and local_only is False:
|
||||
result = await self.redis_cache.async_increment_pipeline(
|
||||
increment_list=increment_list,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
return result
|
||||
except Exception as e:
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.WARNING,
|
||||
"Redis async_increment_cache_pipeline failed; local counters unchanged",
|
||||
"Redis async_increment_cache_pipeline failed, falling back to in-memory result",
|
||||
e,
|
||||
)
|
||||
return None
|
||||
return result
|
||||
|
||||
async def async_set_cache_sadd(self, key, value: list, local_only: bool = False, **kwargs) -> None:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -494,30 +494,6 @@ async def test_async_increment_cache_returns_none_when_no_in_memory_cache_and_re
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_redis_increment_does_not_change_the_local_counter():
|
||||
memory = InMemoryCache()
|
||||
memory.set_cache("counter", 10)
|
||||
redis_cache = MagicMock(spec=RedisCache)
|
||||
redis_cache.async_increment = AsyncMock(side_effect=RuntimeError("redis down"))
|
||||
cache = DualCache(in_memory_cache=memory, redis_cache=redis_cache)
|
||||
|
||||
assert await cache.async_increment_cache("counter", 2) is None
|
||||
assert memory.get_cache("counter") == 10
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_successful_redis_increment_replaces_the_local_counter_with_the_authoritative_value():
|
||||
memory = InMemoryCache()
|
||||
memory.set_cache("counter", 10)
|
||||
redis_cache = MagicMock(spec=RedisCache)
|
||||
redis_cache.async_increment = AsyncMock(return_value=42.0)
|
||||
cache = DualCache(in_memory_cache=memory, redis_cache=redis_cache)
|
||||
|
||||
assert await cache.async_increment_cache("counter", 2) == 42.0
|
||||
assert memory.get_cache("counter") == 42.0
|
||||
|
||||
|
||||
def test_dual_cache_late_attach_redis_wires_writes_and_ttl_sync():
|
||||
"""
|
||||
Typical lazy startup (sync): DualCache runs with in-memory only, then Redis
|
||||
|
|
@ -748,7 +724,7 @@ async def test_redis_timeouts_falling_back_to_memory_log_once_per_interval(caplo
|
|||
assert [(r.levelno, r.getMessage()) for r in visible] == [
|
||||
(
|
||||
logging.WARNING,
|
||||
"Redis async_increment_cache_pipeline failed; local counters unchanged:"
|
||||
"Redis async_increment_cache_pipeline failed, falling back to in-memory result:"
|
||||
" Timeout reading from 127.0.0.1:6379",
|
||||
)
|
||||
]
|
||||
|
|
@ -762,7 +738,7 @@ async def test_redis_timeouts_falling_back_to_memory_log_once_per_interval(caplo
|
|||
assert [(r.levelno, r.getMessage()) for r in caplog.records] == [
|
||||
(
|
||||
logging.WARNING,
|
||||
"Redis async_increment_cache failed; local counter unchanged: Timeout reading from 127.0.0.1:6379"
|
||||
"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)",
|
||||
)
|
||||
]
|
||||
|
|
|
|||
|
|
@ -275,6 +275,45 @@ async def test_native_batch_lookup_and_store_report_partial_hits() -> None:
|
|||
}
|
||||
|
||||
|
||||
async def test_python_batch_callbacks_receive_keys_and_key_value_pairs() -> None:
|
||||
first: Final = object()
|
||||
second: Final = object()
|
||||
|
||||
class CustomCache:
|
||||
def batch_get_cache(self, keys: list[str], *, marker: object) -> tuple[list[str], object]:
|
||||
return keys, marker
|
||||
|
||||
async def async_batch_get_cache(self, keys: list[str], *, marker: object) -> tuple[list[str], object]:
|
||||
return keys, marker
|
||||
|
||||
async def async_set_cache_pipeline(
|
||||
self, cache_list: list[tuple[str, object]], *, marker: object
|
||||
) -> tuple[list[tuple[str, object]], object]:
|
||||
return cache_list, marker
|
||||
|
||||
marker: Final = object()
|
||||
binding: Final = _native.CacheResolver(SimpleNamespace(cache=CustomCache())).resolve()
|
||||
requests: Final = [request("first"), request("second")]
|
||||
|
||||
assert binding.lookup_batch(requests, callback_kwargs={"marker": marker}) == (["first", "second"], marker)
|
||||
assert await binding.async_lookup_batch(requests, callback_kwargs={"marker": marker}) == (
|
||||
["first", "second"],
|
||||
marker,
|
||||
)
|
||||
stored: Final = cast(
|
||||
tuple[list[tuple[str, object]], object],
|
||||
await binding.async_store_batch(
|
||||
requests,
|
||||
[first, second],
|
||||
callback_kwargs={"marker": marker},
|
||||
),
|
||||
)
|
||||
assert [key for key, _ in stored[0]] == ["first", "second"]
|
||||
assert stored[1] is marker
|
||||
assert stored[0][0][1] is first
|
||||
assert stored[0][1][1] is second
|
||||
|
||||
|
||||
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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue