From c233377d9c073fd9e2262c58ea14b2ca75aadf0a Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 21 Sep 2026 22:04:46 +0000 Subject: [PATCH] fix(cache): await valkey semantic batch embeddings inline Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../crates/python-bridge/src/cache/binding.rs | 7 +- .../crates/python-bridge/src/cache/native.rs | 49 ++++- .../python-bridge/src/cache/semantic_step.rs | 175 +++++++++++++----- .../test_valkey_semantic_cache_native.py | 20 ++ 4 files changed, 197 insertions(+), 54 deletions(-) diff --git a/litellm-rust/crates/python-bridge/src/cache/binding.rs b/litellm-rust/crates/python-bridge/src/cache/binding.rs index a8881f19e45..2ff73238202 100644 --- a/litellm-rust/crates/python-bridge/src/cache/binding.rs +++ b/litellm-rust/crates/python-bridge/src/cache/binding.rs @@ -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) diff --git a/litellm-rust/crates/python-bridge/src/cache/native.rs b/litellm-rust/crates/python-bridge/src/cache/native.rs index 7340ef6cce0..260dc8f349e 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native.rs @@ -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> { + 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(), + ), + ) } } } diff --git a/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs b/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs index c62cdb1d9a6..d8f3ab297f8 100644 --- a/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs +++ b/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs @@ -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), } #[derive(Clone, Copy)] @@ -27,9 +26,11 @@ enum State { pub(super) struct SemanticEmbedExecution { backend: Arc>, embedder: PythonEmbedder, - request: ResponseCacheRequest, + requests: Vec>, op: Op, now: Duration, + prepared: Vec>>, + 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>, + embedder: PythonEmbedder, + requests: Vec>, + responses: Vec, + 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 { - 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 { + 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::>()?; - 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( - py: Python<'_>, - cache: Arc>>, - request: ResponseCacheRequest, - op: &Op, - now: Duration, -) -> PyResult { - 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, diff --git a/tests/test_litellm_rust/test_valkey_semantic_cache_native.py b/tests/test_litellm_rust/test_valkey_semantic_cache_native.py index e2bd3dcb10a..fb5c5965333 100644 --- a/tests/test_litellm_rust/test_valkey_semantic_cache_native.py +++ b/tests/test_litellm_rust/test_valkey_semantic_cache_native.py @@ -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]