diff --git a/litellm-rust/crates/cache-response/src/response.rs b/litellm-rust/crates/cache-response/src/response.rs index e70e07a5d26..2f949d511de 100644 --- a/litellm-rust/crates/cache-response/src/response.rs +++ b/litellm-rust/crates/cache-response/src/response.rs @@ -52,6 +52,10 @@ where &self.backend } + pub fn backend_arc(&self) -> &Arc { + &self.backend + } + pub fn default_ttl(&self) -> Option { self.backend.get_ttl(&B::Context::default()) } diff --git a/litellm-rust/crates/cache-valkey-semantic/src/lib.rs b/litellm-rust/crates/cache-valkey-semantic/src/lib.rs index 2bd4bd71bce..e3a3e6094c4 100644 --- a/litellm-rust/crates/cache-valkey-semantic/src/lib.rs +++ b/litellm-rust/crates/cache-valkey-semantic/src/lib.rs @@ -24,6 +24,22 @@ pub trait Embedder: Send + Sync + 'static { ) -> impl Future, Error>> + Send; } +pub struct PreparedEmbedding(pub Vec); + +impl Embedder for PreparedEmbedding { + fn embed(&self, _prompt: &str, _metadata: Option<&Value>) -> Result, Error> { + Ok(self.0.clone()) + } + + async fn async_embed( + &self, + _prompt: &str, + _metadata: Option<&Value>, + ) -> Result, Error> { + Ok(self.0.clone()) + } +} + #[derive(Clone, Debug, PartialEq)] pub struct ValkeySemanticConfig { pub similarity_threshold: f64, @@ -112,6 +128,23 @@ where } } +impl ValkeySemanticCache +where + E: Embedder, + S: CacheCodec + Clone, + C: redis::ConnectionLike + Send + 'static, +{ + pub fn with_embedder(&self, embedder: E2) -> ValkeySemanticCache { + ValkeySemanticCache { + connections: Arc::clone(&self.connections), + embedder, + codec: self.codec.clone(), + config: self.config.clone(), + index_dimension: Arc::clone(&self.index_dimension), + } + } +} + impl BaseCache for ValkeySemanticCache where E: Embedder, @@ -597,8 +630,8 @@ mod tests { use serde_json::{Value, json}; use super::{ - Embedder, ValkeySemanticCache, ValkeySemanticConfig, index_dimension_from_info, - prompt_from_context, scope_tag, + Embedder, PreparedEmbedding, ValkeySemanticCache, ValkeySemanticConfig, + index_dimension_from_info, prompt_from_context, scope_tag, }; #[derive(Clone)] @@ -758,6 +791,47 @@ mod tests { ); } + #[tokio::test] + async fn prepared_embedding_returns_its_vector_for_any_prompt() { + let embedding = PreparedEmbedding(vec![1.0, 2.0]); + assert_eq!( + embedding + .async_embed("different prompt", None) + .await + .unwrap(), + vec![1.0, 2.0] + ); + } + + #[test] + fn with_embedder_shares_index_state_and_connections() { + let entry = CacheEntry { + timestamp: Some(1.0), + response: json!({"answer": "ok"}), + }; + let encoded = ResponseCacheCodec.encode(&entry).unwrap(); + let cache = ValkeySemanticCache::with_connection( + RecordingConnection::new([ok(), ok(), Ok(search_hit(encoded, "0.1"))]), + FixedEmbedder { + vector: vec![1.0, 0.0], + calls: Arc::default(), + }, + ResponseCacheCodec, + ValkeySemanticConfig { + similarity_threshold: 0.8, + index_name: "test".into(), + }, + ); + cache + .set_cache("key", entry.clone(), &semantic_context(None)) + .unwrap(); + let prepared = cache.with_embedder(PreparedEmbedding(vec![1.0, 0.0])); + assert_eq!( + prepared.get_cache("key", &semantic_context(None)).unwrap(), + Some(entry) + ); + } + #[test] fn missing_prompt_does_not_touch_redis() { let cache = ValkeySemanticCache::with_connection( diff --git a/litellm-rust/crates/python-bridge/src/cache/binding.rs b/litellm-rust/crates/python-bridge/src/cache/binding.rs index ad64b24d3c1..a8881f19e45 100644 --- a/litellm-rust/crates/python-bridge/src/cache/binding.rs +++ b/litellm-rust/crates/python-bridge/src/cache/binding.rs @@ -56,12 +56,7 @@ impl ResolvedCache { CacheBinding::Disabled => ready_none(py)?, CacheBinding::Native(service) => { let request = request(input)?; - let service = service.clone(); - run_async( - py, - async move { service.async_lookup(&request, now()).await }, - cache_error, - )? + service.async_lookup_py(py, request)? } CacheBinding::PythonCallback(callback) => callback.async_lookup(py, kwargs)?, }; @@ -179,12 +174,7 @@ impl ResolvedCache { CacheBinding::Native(service) => { let request = self::request(request)?; let response: Value = from_py(response)?; - let service = service.clone(); - run_async( - py, - async move { service.async_store(&request, response, now()).await }, - cache_error, - ) + service.async_store_py(py, request, response) } CacheBinding::PythonCallback(callback) => { callback.async_store(py, response, callback_kwargs) diff --git a/litellm-rust/crates/python-bridge/src/cache/embedder.rs b/litellm-rust/crates/python-bridge/src/cache/embedder.rs index d240d9d019e..3de0ceb3b67 100644 --- a/litellm-rust/crates/python-bridge/src/cache/embedder.rs +++ b/litellm-rust/crates/python-bridge/src/cache/embedder.rs @@ -3,22 +3,37 @@ use std::{future::Future, sync::Arc}; use litellm_cache::Error; use litellm_cache_valkey_semantic::Embedder; use litellm_host_python::to_py; -use pyo3::prelude::*; +use pyo3::{PyTraverseError, PyVisit, prelude::*}; use serde_json::Value; #[derive(Clone)] pub(super) struct PythonEmbedder { sync_embed: Arc>, - async_embed: Arc>, + async_embed_callable: Arc>, } impl PythonEmbedder { pub(super) fn from_backend(backend: &Bound<'_, PyAny>) -> PyResult { Ok(Self { sync_embed: Arc::new(backend.getattr("_get_embedding")?.unbind()), - async_embed: Arc::new(backend.getattr("_get_async_embedding")?.unbind()), + async_embed_callable: Arc::new(backend.getattr("_get_async_embedding")?.unbind()), }) } + + pub(super) fn async_embed_awaitable<'py>( + &self, + py: Python<'py>, + prompt: &str, + metadata: &Option, + ) -> PyResult> { + let metadata = to_py(py, metadata)?; + self.async_embed_callable.bind(py).call1((prompt, metadata)) + } + + pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&*self.sync_embed)?; + visit.call(&*self.async_embed_callable) + } } impl Embedder for PythonEmbedder { @@ -34,25 +49,15 @@ impl Embedder for PythonEmbedder { Ok(result.into_iter().map(|value| value as f32).collect()) } + #[expect( + clippy::manual_async_fn, + reason = "the shared Embedder trait uses an impl Future return" + )] fn async_embed( &self, - prompt: &str, - metadata: Option<&Value>, + _prompt: &str, + _metadata: Option<&Value>, ) -> impl Future, Error>> + Send { - let callable = Arc::clone(&self.async_embed); - let prompt = prompt.to_owned(); - let metadata = metadata.cloned(); - async move { - let future = Python::attach(|py| -> PyResult<_> { - let metadata = to_py(py, &metadata)?; - let awaitable = callable.bind(py).call1((prompt, metadata))?; - pyo3_async_runtimes::tokio::into_future(awaitable) - }) - .map_err(|_| Error::Unavailable)?; - let result = future.await.map_err(|_| Error::Unavailable)?; - let result = Python::attach(|py| result.bind(py).extract::>()) - .map_err(|_| Error::Unavailable)?; - Ok(result.into_iter().map(|value| value as f32).collect()) - } + async { Err(Error::Unavailable) } } } diff --git a/litellm-rust/crates/python-bridge/src/cache/mod.rs b/litellm-rust/crates/python-bridge/src/cache/mod.rs index 4cc87367d91..278d3da1ff9 100644 --- a/litellm-rust/crates/python-bridge/src/cache/mod.rs +++ b/litellm-rust/crates/python-bridge/src/cache/mod.rs @@ -8,6 +8,7 @@ mod handle; mod native; mod request; mod resolver; +mod semantic_step; use litellm_cache::Error; use pyo3::{ diff --git a/litellm-rust/crates/python-bridge/src/cache/native.rs b/litellm-rust/crates/python-bridge/src/cache/native.rs index 3380a914c41..7340ef6cce0 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native.rs @@ -10,9 +10,14 @@ use litellm_cache_response::{ ResponseCacheRequest, WriteBuffer, }; use litellm_cache_valkey_semantic::{ValkeySemanticCache, ValkeySemanticConfig}; +use pyo3::prelude::*; use serde_json::Value; -use super::{embedder::PythonEmbedder, request::NativeRequest}; +use super::{ + embedder::PythonEmbedder, + request::NativeRequest, + semantic_step::{SemanticEmbedExecution, drive_semantic}, +}; fn semantic_key(request: &NativeRequest, scope: &str) -> litellm_cache_response::CacheKeyInput { let mut key = request.key.clone(); @@ -59,6 +64,7 @@ pub(super) enum NativeResponseCache { }, ValkeySemantic { cache: Arc>>, + embedder: PythonEmbedder, scope: String, }, } @@ -100,7 +106,7 @@ impl NativeResponseCache { ) -> Result { let backend = ValkeySemanticCache::new( url, - embedder, + embedder.clone(), ResponseCacheCodec, ValkeySemanticConfig { similarity_threshold, @@ -109,6 +115,7 @@ impl NativeResponseCache { )?; Ok(Self::ValkeySemantic { cache: Arc::new(ResponseCache::new(Arc::new(backend))), + embedder, scope: String::from("key"), }) } @@ -152,7 +159,13 @@ impl NativeResponseCache { pub fn with_scope(self, scope: String) -> Self { match self { - Self::ValkeySemantic { cache, .. } => Self::ValkeySemantic { cache, scope }, + Self::ValkeySemantic { + cache, embedder, .. + } => Self::ValkeySemantic { + cache, + embedder, + scope, + }, value => value, } } @@ -215,7 +228,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.lookup(&Self::exact(request), now), Self::Redis { cache, .. } => cache.lookup(&Self::exact(request), now), - Self::ValkeySemantic { cache, scope } => { + Self::ValkeySemantic { cache, scope, .. } => { cache.lookup(&Self::semantic(request, scope), now) } } @@ -230,7 +243,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.store(&Self::exact(request), response, now), Self::Redis { cache, .. } => cache.store(&Self::exact(request), response, now), - Self::ValkeySemantic { cache, scope } => { + Self::ValkeySemantic { cache, scope, .. } => { cache.store(&Self::semantic(request, scope), response, now) } } @@ -262,7 +275,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.async_lookup(&Self::exact(request), now).await, Self::Redis { cache, .. } => cache.async_lookup(&Self::exact(request), now).await, - Self::ValkeySemantic { cache, scope } => { + Self::ValkeySemantic { cache, scope, .. } => { cache .async_lookup(&Self::semantic(request, scope), now) .await @@ -270,6 +283,36 @@ impl NativeResponseCache { } } + pub(super) fn async_lookup_py<'py>( + &self, + py: Python<'py>, + request: NativeRequest, + ) -> PyResult> { + match self { + Self::Memory(_) | Self::Redis { .. } => { + let service = self.clone(); + litellm_host_python::run_async( + py, + async move { service.async_lookup(&request, super::request::now()).await }, + super::cache_error, + ) + } + Self::ValkeySemantic { + cache, + embedder, + scope, + } => drive_semantic( + py, + SemanticEmbedExecution::lookup( + Arc::clone(cache.backend_arc()), + embedder.clone(), + Self::semantic(&request, scope), + super::request::now(), + ), + ), + } + } + pub async fn async_store( &self, request: &NativeRequest, @@ -298,7 +341,7 @@ impl NativeResponseCache { .async_store(cache, &Self::exact(request), response, now) .await } - Self::ValkeySemantic { cache, scope } => { + Self::ValkeySemantic { cache, scope, .. } => { cache .async_store(&Self::semantic(request, scope), response, now) .await @@ -306,6 +349,42 @@ impl NativeResponseCache { } } + pub(super) fn async_store_py<'py>( + &self, + py: Python<'py>, + request: NativeRequest, + response: Value, + ) -> PyResult> { + match self { + Self::Memory(_) | Self::Redis { .. } => { + let service = self.clone(); + litellm_host_python::run_async( + py, + async move { + service + .async_store(&request, response, super::request::now()) + .await + }, + super::cache_error, + ) + } + Self::ValkeySemantic { + cache, + embedder, + scope, + } => drive_semantic( + py, + SemanticEmbedExecution::store( + Arc::clone(cache.backend_arc()), + embedder.clone(), + Self::semantic(&request, scope), + response, + super::request::now(), + ), + ), + } + } + pub async fn async_lookup_batch( &self, requests: &[NativeRequest], @@ -344,12 +423,10 @@ impl NativeResponseCache { .collect(); cache.async_store_batch(entries, now).await } - Self::ValkeySemantic { cache, scope } => { - let entries = entries - .into_iter() - .map(|(request, value)| (Self::semantic(&request, scope), value)) - .collect(); - 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) + }) } } } diff --git a/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs b/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs new file mode 100644 index 00000000000..c62cdb1d9a6 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs @@ -0,0 +1,154 @@ +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_host_python::{Execution, ExecutionBody, ExecutionStep, run_async}; +use pyo3::{PyTraverseError, PyVisit, exceptions::PyRuntimeError, prelude::*}; +use serde_json::Value; + +use super::{cache_error, embedder::PythonEmbedder}; + +pub(super) enum Op { + Lookup, + Store(Value), +} + +#[derive(Clone, Copy)] +enum State { + Start, + AwaitingEmbedding, + AwaitingStorage, + Done, +} + +pub(super) struct SemanticEmbedExecution { + backend: Arc>, + embedder: PythonEmbedder, + request: ResponseCacheRequest, + op: Op, + now: Duration, + state: State, +} + +impl SemanticEmbedExecution { + pub(super) fn lookup( + backend: Arc>, + embedder: PythonEmbedder, + request: ResponseCacheRequest, + now: Duration, + ) -> Self { + Self { + backend, + embedder, + request, + op: Op::Lookup, + now, + state: State::Start, + } + } + + pub(super) fn store( + backend: Arc>, + embedder: PythonEmbedder, + request: ResponseCacheRequest, + response: Value, + now: Duration, + ) -> Self { + Self { + backend, + embedder, + request, + op: Op::Store(response), + now, + 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); + }; + let awaitable = + self.embedder + .async_embed_awaitable(py, &prompt, &self.request.context.metadata)?; + self.state = State::AwaitingEmbedding; + Ok(ExecutionStep::Await(awaitable.unbind())) + } + + fn resume_py( + &mut self, + py: Python<'_>, + result: Option>>, + ) -> PyResult { + match (self.state, result) { + (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) + } + (State::AwaitingStorage, Some(Ok(value))) => { + self.state = State::Done; + Ok(ExecutionStep::Return(value)) + } + (_, Some(Err(error))) => Err(error), + _ => Err(PyRuntimeError::new_err( + "invalid semantic cache execution state", + )), + } + } +} + +impl ExecutionBody for SemanticEmbedExecution { + fn resume(&mut self, result: Option>>) -> PyResult { + Python::attach(|py| self.resume_py(py, result)) + } + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + self.embedder.traverse(visit) + } +} + +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, +) -> PyResult> { + let execution = Py::new(py, Execution::new(body))?; + py.import("litellm.rust_bridge.lifecycle")? + .getattr("drive")? + .call1((execution,)) +} 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 81fcf00ffc0..e2bd3dcb10a 100644 --- a/tests/test_litellm_rust/test_valkey_semantic_cache_native.py +++ b/tests/test_litellm_rust/test_valkey_semantic_cache_native.py @@ -1,6 +1,9 @@ +import asyncio +import contextvars import hashlib import os import struct +import threading import time from collections.abc import Generator, Mapping from types import SimpleNamespace @@ -16,6 +19,7 @@ from litellm.rust_bridge import _native from litellm.types.caching import LiteLLMCacheType pytestmark: Final = pytest.mark.requires_rust_extension +embedding_context: Final = contextvars.ContextVar("embedding_context") @pytest.fixture @@ -171,6 +175,43 @@ async def test_async_lookup_and_store( assert await binding.async_lookup(request) == {"answer": "async"} +async def test_async_embedding_runs_inline_in_caller_task( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + observed: dict[str, object] = {} + + async def async_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]: + observed["context"] = embedding_context.get("missing") + observed["task"] = asyncio.current_task() + observed["thread"] = threading.get_ident() + embedding_context.set("embedder") + return [1.0, 0.0] + + 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() + request: Final = {**_request(), "ttl_seconds": 2.0} + caller_task: Final = asyncio.current_task() + caller_thread: Final = threading.get_ident() + token: Final = embedding_context.set("caller") + try: + await binding.async_store(request, {"answer": "inline"}) + assert observed["context"] == "caller" + assert observed["task"] is caller_task + assert observed["thread"] == caller_thread + assert embedding_context.get() == "embedder" + assert await binding.async_lookup(request) == {"answer": "inline"} + finally: + embedding_context.reset(token) + + def test_facade_activation_and_mutation_fallback( valkey_url: str, index_name: str,