mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
fix(cache): await valkey semantic batch embeddings inline
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
847f732f5d
commit
c233377d9c
4 changed files with 197 additions and 54 deletions
|
|
@ -231,12 +231,7 @@ impl ResolvedCache {
|
|||
));
|
||||
}
|
||||
let entries = requests.into_iter().zip(responses).collect();
|
||||
let service = service.clone();
|
||||
run_async(
|
||||
py,
|
||||
async move { service.async_store_batch(entries, now()).await },
|
||||
cache_error,
|
||||
)
|
||||
service.async_store_batch_py(py, entries)
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => {
|
||||
callback.async_store_batch(py, callback_result, callback_kwargs)
|
||||
|
|
|
|||
|
|
@ -424,9 +424,52 @@ impl NativeResponseCache {
|
|||
cache.async_store_batch(entries, now).await
|
||||
}
|
||||
Self::ValkeySemantic { cache, scope, .. } => {
|
||||
entries.into_iter().try_for_each(|(request, value)| {
|
||||
cache.store(&Self::semantic(&request, scope), value, now)
|
||||
})
|
||||
let entries = entries
|
||||
.into_iter()
|
||||
.map(|(request, value)| (Self::semantic(&request, scope), value))
|
||||
.collect();
|
||||
cache.async_store_batch(entries, now).await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn async_store_batch_py<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
entries: Vec<(NativeRequest, Value)>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
match self {
|
||||
Self::Memory(_) | Self::Redis { .. } => {
|
||||
let service = self.clone();
|
||||
litellm_host_python::run_async(
|
||||
py,
|
||||
async move {
|
||||
service
|
||||
.async_store_batch(entries, super::request::now())
|
||||
.await
|
||||
},
|
||||
super::cache_error,
|
||||
)
|
||||
}
|
||||
Self::ValkeySemantic {
|
||||
cache,
|
||||
embedder,
|
||||
scope,
|
||||
} => {
|
||||
let (requests, responses): (Vec<_>, Vec<_>) = entries
|
||||
.into_iter()
|
||||
.map(|(request, response)| (Self::semantic(&request, scope), response))
|
||||
.unzip();
|
||||
drive_semantic(
|
||||
py,
|
||||
SemanticEmbedExecution::store_batch(
|
||||
Arc::clone(cache.backend_arc()),
|
||||
embedder.clone(),
|
||||
requests,
|
||||
responses,
|
||||
super::request::now(),
|
||||
),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,9 +2,7 @@ use std::{sync::Arc, time::Duration};
|
|||
|
||||
use litellm_cache::SemanticCacheContext;
|
||||
use litellm_cache_response::{ResponseCache, ResponseCacheCodec, ResponseCacheRequest};
|
||||
use litellm_cache_valkey_semantic::{
|
||||
Embedder, PreparedEmbedding, ValkeySemanticCache, prompt_from_context,
|
||||
};
|
||||
use litellm_cache_valkey_semantic::{PreparedEmbedding, ValkeySemanticCache, prompt_from_context};
|
||||
use litellm_host_python::{Execution, ExecutionBody, ExecutionStep, run_async};
|
||||
use pyo3::{PyTraverseError, PyVisit, exceptions::PyRuntimeError, prelude::*};
|
||||
use serde_json::Value;
|
||||
|
|
@ -14,6 +12,7 @@ use super::{cache_error, embedder::PythonEmbedder};
|
|||
pub(super) enum Op {
|
||||
Lookup,
|
||||
Store(Value),
|
||||
StoreBatch(Vec<Value>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
|
|
@ -27,9 +26,11 @@ enum State {
|
|||
pub(super) struct SemanticEmbedExecution {
|
||||
backend: Arc<ValkeySemanticCache<PythonEmbedder, ResponseCacheCodec>>,
|
||||
embedder: PythonEmbedder,
|
||||
request: ResponseCacheRequest<SemanticCacheContext>,
|
||||
requests: Vec<ResponseCacheRequest<SemanticCacheContext>>,
|
||||
op: Op,
|
||||
now: Duration,
|
||||
prepared: Vec<Option<Vec<f32>>>,
|
||||
index: usize,
|
||||
state: State,
|
||||
}
|
||||
|
||||
|
|
@ -43,9 +44,11 @@ impl SemanticEmbedExecution {
|
|||
Self {
|
||||
backend,
|
||||
embedder,
|
||||
request,
|
||||
requests: vec![request],
|
||||
op: Op::Lookup,
|
||||
now,
|
||||
prepared: vec![None],
|
||||
index: 0,
|
||||
state: State::Start,
|
||||
}
|
||||
}
|
||||
|
|
@ -60,23 +63,132 @@ impl SemanticEmbedExecution {
|
|||
Self {
|
||||
backend,
|
||||
embedder,
|
||||
request,
|
||||
requests: vec![request],
|
||||
op: Op::Store(response),
|
||||
now,
|
||||
prepared: vec![None],
|
||||
index: 0,
|
||||
state: State::Start,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn store_batch(
|
||||
backend: Arc<ValkeySemanticCache<PythonEmbedder, ResponseCacheCodec>>,
|
||||
embedder: PythonEmbedder,
|
||||
requests: Vec<ResponseCacheRequest<SemanticCacheContext>>,
|
||||
responses: Vec<Value>,
|
||||
now: Duration,
|
||||
) -> Self {
|
||||
Self {
|
||||
backend,
|
||||
embedder,
|
||||
prepared: vec![None; requests.len()],
|
||||
requests,
|
||||
op: Op::StoreBatch(responses),
|
||||
now,
|
||||
index: 0,
|
||||
state: State::Start,
|
||||
}
|
||||
}
|
||||
|
||||
fn start(&mut self, py: Python<'_>) -> PyResult<ExecutionStep> {
|
||||
let Some(prompt) = prompt_from_context(&self.request.context) else {
|
||||
let cache = Arc::new(ResponseCache::new(Arc::clone(&self.backend)));
|
||||
self.state = State::AwaitingStorage;
|
||||
return storage_step(py, cache, self.request.clone(), &self.op, self.now);
|
||||
while self.index < self.requests.len() {
|
||||
let request = &self.requests[self.index];
|
||||
let Some(prompt) = prompt_from_context(&request.context) else {
|
||||
self.index += 1;
|
||||
continue;
|
||||
};
|
||||
let metadata = request.context.metadata.clone();
|
||||
let awaitable = self
|
||||
.embedder
|
||||
.async_embed_awaitable(py, &prompt, &metadata)?;
|
||||
self.state = State::AwaitingEmbedding;
|
||||
return Ok(ExecutionStep::Await(awaitable.unbind()));
|
||||
}
|
||||
self.state = State::AwaitingStorage;
|
||||
self.storage_step(py)
|
||||
}
|
||||
|
||||
fn storage_step(&self, py: Python<'_>) -> PyResult<ExecutionStep> {
|
||||
let requests = self.requests.clone();
|
||||
let prepared = self.prepared.clone();
|
||||
let backend = Arc::clone(&self.backend);
|
||||
let now = self.now;
|
||||
let awaitable = match &self.op {
|
||||
Op::Lookup => {
|
||||
let Some(request) = requests.into_iter().next() else {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
"semantic lookup requires one request",
|
||||
));
|
||||
};
|
||||
match prepared.into_iter().next().flatten() {
|
||||
Some(values) => {
|
||||
let backend = backend.with_embedder(PreparedEmbedding(values));
|
||||
let cache = Arc::new(ResponseCache::new(Arc::new(backend)));
|
||||
run_async(
|
||||
py,
|
||||
async move { cache.async_lookup(&request, now).await },
|
||||
cache_error,
|
||||
)?
|
||||
}
|
||||
None => {
|
||||
let cache = Arc::new(ResponseCache::new(backend));
|
||||
run_async(
|
||||
py,
|
||||
async move { cache.async_lookup(&request, now).await },
|
||||
cache_error,
|
||||
)?
|
||||
}
|
||||
}
|
||||
}
|
||||
Op::Store(response) => {
|
||||
let Some(request) = requests.into_iter().next() else {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
"semantic store requires one request",
|
||||
));
|
||||
};
|
||||
let response = response.clone();
|
||||
match prepared.into_iter().next().flatten() {
|
||||
Some(values) => {
|
||||
let backend = backend.with_embedder(PreparedEmbedding(values));
|
||||
let cache = Arc::new(ResponseCache::new(Arc::new(backend)));
|
||||
run_async(
|
||||
py,
|
||||
async move { cache.async_store(&request, response, now).await },
|
||||
cache_error,
|
||||
)?
|
||||
}
|
||||
None => {
|
||||
let cache = Arc::new(ResponseCache::new(backend));
|
||||
run_async(
|
||||
py,
|
||||
async move { cache.async_store(&request, response, now).await },
|
||||
cache_error,
|
||||
)?
|
||||
}
|
||||
}
|
||||
}
|
||||
Op::StoreBatch(responses) => {
|
||||
let responses = responses.clone();
|
||||
run_async(
|
||||
py,
|
||||
async move {
|
||||
for ((request, response), prepared) in
|
||||
requests.into_iter().zip(responses).zip(prepared)
|
||||
{
|
||||
let Some(values) = prepared else {
|
||||
continue;
|
||||
};
|
||||
let backend = backend.with_embedder(PreparedEmbedding(values));
|
||||
let cache = ResponseCache::new(Arc::new(backend));
|
||||
cache.async_store(&request, response, now).await?;
|
||||
}
|
||||
Ok(())
|
||||
},
|
||||
cache_error,
|
||||
)?
|
||||
}
|
||||
};
|
||||
let awaitable =
|
||||
self.embedder
|
||||
.async_embed_awaitable(py, &prompt, &self.request.context.metadata)?;
|
||||
self.state = State::AwaitingEmbedding;
|
||||
Ok(ExecutionStep::Await(awaitable.unbind()))
|
||||
}
|
||||
|
||||
|
|
@ -89,12 +201,10 @@ impl SemanticEmbedExecution {
|
|||
(State::Start, None) => self.start(py),
|
||||
(State::AwaitingEmbedding, Some(Ok(value))) => {
|
||||
let values = value.bind(py).extract::<Vec<f64>>()?;
|
||||
let backend = self.backend.with_embedder(PreparedEmbedding(
|
||||
values.into_iter().map(|value| value as f32).collect(),
|
||||
));
|
||||
let cache = Arc::new(ResponseCache::new(Arc::new(backend)));
|
||||
self.state = State::AwaitingStorage;
|
||||
storage_step(py, cache, self.request.clone(), &self.op, self.now)
|
||||
self.prepared[self.index] =
|
||||
Some(values.into_iter().map(|value| value as f32).collect());
|
||||
self.index += 1;
|
||||
self.start(py)
|
||||
}
|
||||
(State::AwaitingStorage, Some(Ok(value))) => {
|
||||
self.state = State::Done;
|
||||
|
|
@ -118,31 +228,6 @@ impl ExecutionBody for SemanticEmbedExecution {
|
|||
}
|
||||
}
|
||||
|
||||
fn storage_step<E: Embedder>(
|
||||
py: Python<'_>,
|
||||
cache: Arc<ResponseCache<ValkeySemanticCache<E, ResponseCacheCodec>>>,
|
||||
request: ResponseCacheRequest<SemanticCacheContext>,
|
||||
op: &Op,
|
||||
now: Duration,
|
||||
) -> PyResult<ExecutionStep> {
|
||||
let awaitable = match op {
|
||||
Op::Lookup => run_async(
|
||||
py,
|
||||
async move { cache.async_lookup(&request, now).await },
|
||||
cache_error,
|
||||
)?,
|
||||
Op::Store(response) => {
|
||||
let response = response.clone();
|
||||
run_async(
|
||||
py,
|
||||
async move { cache.async_store(&request, response, now).await },
|
||||
cache_error,
|
||||
)?
|
||||
}
|
||||
};
|
||||
Ok(ExecutionStep::Await(awaitable.unbind()))
|
||||
}
|
||||
|
||||
pub(super) fn drive_semantic<'py>(
|
||||
py: Python<'py>,
|
||||
body: SemanticEmbedExecution,
|
||||
|
|
|
|||
|
|
@ -332,11 +332,31 @@ async def test_async_store_batch_and_lookup(
|
|||
index_name,
|
||||
{"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]},
|
||||
)
|
||||
sync_calls: Final = []
|
||||
async_tasks: Final = []
|
||||
|
||||
def sync_embedding(prompt: str, metadata: Mapping[str, object] | None = None) -> list[float]:
|
||||
sync_calls.append(prompt)
|
||||
return {"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]}[prompt]
|
||||
|
||||
async def async_embedding(
|
||||
prompt: str,
|
||||
metadata: dict[str, object] | None = None,
|
||||
) -> list[float]:
|
||||
async_tasks.append(asyncio.current_task())
|
||||
return {"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]}[prompt]
|
||||
|
||||
backend._get_embedding = sync_embedding
|
||||
backend._get_async_embedding = async_embedding
|
||||
handle: Final = _native._CacheTestHandle.valkey_semantic(valkey_url, 0.8, index_name, backend)
|
||||
binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve()
|
||||
requests: Final = [_request("prompt A"), _request("prompt B")]
|
||||
responses: Final = [{"answer": "A"}, {"answer": "B"}]
|
||||
caller_task: Final = asyncio.current_task()
|
||||
await binding.async_store_batch(requests, responses)
|
||||
assert sync_calls == []
|
||||
assert async_tasks
|
||||
assert all(task is caller_task for task in async_tasks)
|
||||
assert await binding.async_lookup(requests[0]) == responses[0]
|
||||
assert await binding.async_lookup(requests[1]) == responses[1]
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue