refactor(cache): remove dead Cache._native_cache runtime path (#44785)

The native response-cache runtime attached through Cache._native_cache is
unreachable since the V2 cache replaced it. Drop the Python branches and
wrapper, the _ResponseCacheRuntime pyclass and its backend/activation/
semantic modules, the python-bridge deps only they used, and the tests and
fixtures dedicated to that path.

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-05 20:58:19 -07:00 • committed by GitHub
parent a983b2a5e7
commit 5f9eff55f8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
37 changed files with 28 additions and 7875 deletions

View file

@ -4112,16 +4112,9 @@ dependencies = [
"litellm-auth",
"litellm-auth-aws",
"litellm-cache",
"litellm-cache-azure-blob",
"litellm-cache-disk",
"litellm-cache-gcs",
"litellm-cache-memory",
"litellm-cache-qdrant-semantic",
"litellm-cache-redis",
"litellm-cache-redis-semantic",
"litellm-cache-response",
"litellm-cache-s3",
"litellm-cache-valkey-semantic",
"litellm-callbacks-legacy-python",
"litellm-core",
"litellm-core-utils",
@ -4142,21 +4135,18 @@ dependencies = [
"prost",
"pyo3",
"pyo3-async-runtimes",
"qdrant-client",
"redis",
"reqwest 0.12.28",
"rstest",
"serde",
"serde_json",
"serde_with",
"sha2 0.10.9",
"sqlx",
"strum",
"thiserror 2.0.19",
"tokio",
"tokio-tungstenite",
"tracing",
"url",
"veil",
"wiremock",
]

View file

@ -29,17 +29,9 @@ litellm-host.workspace = true
bytes.workspace = true
futures-util.workspace = true
litellm-cache.workspace = true
litellm-cache-azure-blob.workspace = true
litellm-cache-memory.workspace = true
litellm-cache-redis.workspace = true
litellm-cache-s3.workspace = true
litellm-cache-gcs.workspace = true
litellm-cache-disk.workspace = true
litellm-cache-redis-semantic.workspace = true
litellm-cache-response.workspace = true
litellm-cache-qdrant-semantic.workspace = true
qdrant-client.workspace = true
litellm-cache-valkey-semantic = { path = "../cache-valkey-semantic" }
serde.workspace = true
litellm-auth.workspace = true
litellm-auth-aws.workspace = true
@ -63,7 +55,6 @@ strum.workspace = true
veil.workspace = true
thiserror.workspace = true
tokio = { workspace = true, features = ["rt", "sync"] }
url.workspace = true
[dev-dependencies]
tracing.workspace = true
@ -75,7 +66,6 @@ criterion.workspace = true
futures-util.workspace = true
rstest.workspace = true
sqlx = { workspace = true, features = ["migrate"] }
sha2.workspace = true
tokio-tungstenite.workspace = true
wiremock.workspace = true
aws-sdk-secretsmanager = "1.117.0"

View file

@ -1,13 +1,7 @@
# Cache boundary
This folder owns Python cache API compatibility: argument projection, facade identity, public result construction, Python embedding calls and per-operation composition of native backends. Cache algorithms, storage protocols and response-cache semantics belong to their cache crates
This folder owns how Rust inference reaches the selected cache: global cache selection, route admission, delegation to a Python cache and the experimental V2 native handles. Cache algorithms, storage protocols and response-cache semantics belong to their cache crates
`mod.rs` exposes the cache boundary to routes and module registration; adapter directories remain private. `selection.rs` owns global cache selection, route admission and inference protocol composition for both adapters. `runtime.rs` exposes the Python-facing runtime that can wrap either native storage or a Python callback. `future.rs` converts cache results into ready Futures
`mod.rs` exposes the cache boundary to routes and module registration; adapter directories remain private. `selection.rs` owns global cache selection, route admission and inference protocol composition for both adapters
`python/` delegates operations to the selected Python cache without discovering configuration. `native/` owns native backend construction, configuration projection, facade validation, embedding and storage bindings, including experimental V2 handles. Neither adapter depends on shared selection or the other adapter. Shared composition depends on the adapters, and routes use only the parent module's exports
`SemanticExecution` belongs here because its steps select cache operations and invoke the Python embedder. Use the shared `Execution` handle and inline lifecycle driver; do not duplicate coroutine state validation, runtime waiting or GIL machinery. Python embedding awaits stay in the caller's task, and cancellation must prevent later backend or batch operations from starting
Resolved asyncio Future construction is generic host machinery. Use `litellm-host-python::ready_future` with an already constructed Python value. Keep cache-specific conversion and disabled-cache return values here. Preserve the Future-returning API and running-loop requirement
Tests for Future mechanics belong in `host-python`; tests for disabled-cache values, embedding failure policy, batch sequencing and cancellation belong with this cache adapter. Assert public behavior rather than the location or name of a helper
`python/` delegates operations to the selected Python cache without discovering configuration. `native/` owns the experimental V2 native cache handles. Neither adapter depends on shared selection or the other adapter. Shared composition depends on the adapters, and routes use only the parent module's exports

View file

@ -1,13 +0,0 @@
use litellm_host_python::{ready_future, to_py};
use pyo3::prelude::*;
pub(super) fn ready_none(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
ready_value(py, &())
}
pub(super) fn ready_value<'py, T: serde::Serialize>(
py: Python<'py>,
value: &T,
) -> PyResult<Bound<'py, PyAny>> {
ready_future(py, to_py(py, value)?.bind(py))
}

View file

@ -1,12 +1,9 @@
mod future;
mod native;
mod python;
mod runtime;
mod selection;
pub(crate) use native::NativeCacheHandle;
pub(crate) use python::{CacheCall, PythonCache};
pub(crate) use runtime::ResolvedCache;
pub(crate) use selection::{Cached, Selection, admit_native, configure, configured_native};
use litellm_cache::Error;

View file

@ -1,9 +1,9 @@
# Native cache bindings
This directory constructs and exposes Rust cache backends to Python. It owns backend configuration projection, facade validation, native request conversion, semantic embedding integration and experimental V2 handles. Cache algorithms and storage protocols remain in their cache crates
This directory exposes the experimental V2 native cache handles (`NativeCacheHandle`) to Python. Cache algorithms and storage protocols remain in their cache crates
Accept the cache object or projected configuration selected by the parent module. Do not read global `litellm.cache`, decide route admission, or select the Python cache adapter here
Accept the cache object selected by the parent module. Do not read global `litellm.cache`, decide route admission, or select the Python cache adapter here
Keep Python-facing cache classes and method signatures stable when reorganizing modules. Native internals stay private to this directory unless the shared cache boundary or Python module registration needs them. Python embedding awaits use the existing inline lifecycle driver, preserving caller task identity and cancellation
Keep Python-facing cache classes and method signatures stable when reorganizing modules. Native internals stay private to this directory unless the shared cache boundary or Python module registration needs them
Verify changes with the existing backend and facade tests using a freshly built extension. Test cache behavior, not module paths or file structure
Verify changes with the V2 cache tests using a freshly built extension. Test cache behavior, not module paths or file structure

View file

@ -1,112 +0,0 @@
use crate::cache::cache_error;
use crate::execution::run_sync_value;
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
use litellm_cache_redis_semantic::RedisSemanticConfig;
use litellm_host_python::release_gil;
use litellm_http::ClientVariant;
use pyo3::prelude::*;
use super::{
backend::NativeResponseCache,
config::{CacheBackendConfig, NativeCacheConfig, UnsupportedCacheConfig},
embedder::PythonEmbedder,
};
use crate::errors::RustBridgeDeclined;
use crate::http::host_client;
fn declined(reason: UnsupportedCacheConfig) -> PyErr {
RustBridgeDeclined::new_err(reason.message())
}
/// Builds the native backend a `Cache` facade's projected configuration describes. `backend` is
/// the facade's `.cache` object, which owns embedding for the Python-embedded semantic caches.
pub(in crate::cache) fn activate(
py: Python<'_>,
backend: &Bound<'_, PyAny>,
config: NativeCacheConfig,
) -> PyResult<NativeResponseCache> {
let policy = config.policy;
let service = match config.backend {
CacheBackendConfig::Memory(memory) => {
NativeResponseCache::memory(memory.capacity, memory.default_ttl, memory.max_entry_bytes)
}
CacheBackendConfig::Redis(redis) => {
let url = redis.connection.native_url().map_err(declined)?;
let flush_size = policy.redis_flush_size.map(|_| redis.flush_size);
release_gil(py, move || {
NativeResponseCache::redis(
&url,
&redis.topology,
Some(redis.default_ttl),
redis.namespace,
)
})
.map_err(cache_error)?
.with_redis_flush_size(flush_size)
}
CacheBackendConfig::S3(s3) => {
let http = host_client(py, ClientVariant::NoRedirect)?;
run_sync_value(
py,
async move { Ok(NativeResponseCache::s3(*s3, http).await) },
)?
}
CacheBackendConfig::Gcs(gcs) => NativeResponseCache::gcs(
GcsConfig {
bucket_name: gcs.bucket_name,
gcs_path: Some(gcs.key_prefix),
path_service_account: gcs.path_service_account,
endpoint: DEFAULT_ENDPOINT.to_owned(),
},
host_client(py, ClientVariant::NoRedirect)?,
None,
),
CacheBackendConfig::Disk(disk) => {
release_gil(py, move || NativeResponseCache::disk(&disk.directory))
.map_err(cache_error)?
}
CacheBackendConfig::AzureBlob(azure) => {
let http = host_client(py, ClientVariant::NoRedirect)?;
run_sync_value(py, async move {
NativeResponseCache::azure_blob(&azure.account_url, &azure.container, http)
.await
.map_err(cache_error)
})?
}
CacheBackendConfig::RedisSemantic(semantic) => {
let url = semantic.native_url().map_err(declined)?.to_owned();
let embedder = PythonEmbedder::new(backend.clone().unbind());
let semantic_config = RedisSemanticConfig {
index_name: semantic.index_name,
similarity_threshold: semantic.similarity_threshold as f32,
};
release_gil(py, move || {
NativeResponseCache::redis_semantic(&url, embedder, semantic_config)
})
.map_err(cache_error)?
}
CacheBackendConfig::ValkeySemantic(valkey) => {
let url = valkey.connection.native_url().map_err(declined)?;
let embedder = PythonEmbedder::new(backend.clone().unbind());
release_gil(py, move || {
NativeResponseCache::valkey_semantic(
&url,
valkey.similarity_threshold,
valkey.index_name,
embedder,
)
})
.map_err(cache_error)?
}
CacheBackendConfig::QdrantSemantic(qdrant) => {
let client = host_client(py, ClientVariant::Provider)?;
run_sync_value(py, async move {
let runtime = tokio::runtime::Handle::current();
NativeResponseCache::qdrant_semantic(*qdrant, client, runtime)
.await
.map_err(cache_error)
})?
}
};
Ok(service.with_scope(policy.semantic_cache_scope))
}

View file

@ -1,690 +0,0 @@
use crate::cache::cache_error;
use std::{sync::Arc, time::Duration};
use litellm_cache::{CacheCodec, CacheConnectionResult, Error, semantic::SemanticLookup};
use litellm_cache_azure_blob::AzureBlobCache;
use litellm_cache_disk::DiskCache;
use litellm_cache_gcs::{GcsCache, GcsConfig, StaticTokenSource};
use litellm_cache_memory::InMemoryCache;
use litellm_cache_qdrant_semantic::{OpenAiEmbedder, QdrantSemanticCache};
use litellm_cache_redis::{RedisCache, RedisTopology};
use litellm_cache_redis_semantic::{RedisSemanticCache, RedisSemanticConfig};
use litellm_cache_response::{
ConnectionProbe, ExactResponseCache, PartialHits, ResponseCache, ResponseCacheCodec,
WriteBuffer,
};
use litellm_cache_s3::{S3Cache, S3CacheConfig};
use litellm_cache_valkey_semantic::{ValkeySemanticCache, ValkeySemanticConfig};
use pyo3::prelude::*;
use serde_json::Value;
use super::{
config::QdrantSemanticCacheConfig,
embedder::PythonEmbedder,
identity::BackendIdentity,
request::{NativeRequest, now},
semantic::{EmbeddingFailure, SemanticExecution, SemanticOperation, drive},
};
/// What the Python embedder receives for one semantic request.
pub(in crate::cache) struct EmbeddingInput {
pub(in crate::cache) prompt: String,
pub(in crate::cache) metadata: Option<Value>,
}
/// An exact-match backend behind one pointer, with the identity its facade must reproduce.
pub(in crate::cache) struct ExactService {
cache: Arc<dyn ExactResponseCache>,
probe: Option<Arc<dyn ConnectionProbe>>,
buffer: Option<WriteBuffer>,
identity: BackendIdentity,
}
#[derive(Clone)]
pub(in crate::cache) enum NativeResponseCache {
Exact(Arc<ExactService>),
ValkeySemantic {
cache: Arc<ResponseCache<ValkeySemanticCache<PythonEmbedder, ResponseCacheCodec>>>,
embedder: PythonEmbedder,
scope: String,
},
RedisSemantic {
cache: Arc<ResponseCache<RedisSemanticCache<PythonEmbedder, ResponseCacheCodec>>>,
embedder: PythonEmbedder,
},
QdrantSemantic(Arc<ResponseCache<QdrantSemanticCache<OpenAiEmbedder, ResponseCacheCodec>>>),
}
impl NativeResponseCache {
pub fn memory(capacity: usize, ttl: Duration, max_entry_bytes: usize) -> Self {
let backend = InMemoryCache::with_clock_and_size_measurement(
Some(capacity),
Some(ttl),
Some(max_entry_bytes),
Some(Arc::new(|entry| {
ResponseCacheCodec.encode(entry).map(|bytes| bytes.len())
})),
now,
);
let identity = BackendIdentity::Memory {
capacity: backend.max_size_in_memory(),
max_entry_bytes: backend.max_entry_bytes(),
default_ttl: None,
};
Self::exact(ResponseCache::new(Arc::new(backend)), identity)
}
pub fn redis(
url: &str,
topology: &RedisTopology,
ttl: Option<Duration>,
namespace: Option<String>,
) -> Result<Self, Error> {
let backend =
RedisCache::connect(url, topology, ttl, ResponseCacheCodec)?.with_namespace(namespace);
let identity = BackendIdentity::Redis {
topology: backend.topology().clone(),
namespace: backend.namespace().map(str::to_owned),
default_ttl: None,
};
Ok(Self::exact_probed(
ResponseCache::new(Arc::new(backend)),
identity,
))
}
pub async fn s3(config: S3CacheConfig, http: litellm_http::Client) -> Self {
let runtime = tokio::runtime::Handle::current();
let backend = S3Cache::new(config, http, ResponseCacheCodec, runtime);
let identity = BackendIdentity::S3 {
bucket: backend.bucket().to_owned(),
key_prefix: backend.key_prefix().to_owned(),
region: backend.region().to_owned(),
endpoint: backend.endpoint().map(str::to_owned),
};
Self::exact(ResponseCache::new(Arc::new(backend)), identity)
}
pub fn disk(directory: impl AsRef<std::path::Path>) -> Result<Self, Error> {
let backend = DiskCache::open(directory, ResponseCacheCodec)?;
let identity = BackendIdentity::Disk {
directory: backend.directory().to_path_buf(),
};
Ok(Self::exact(ResponseCache::new(Arc::new(backend)), identity))
}
pub fn gcs(config: GcsConfig, client: litellm_http::Client, token: Option<String>) -> Self {
let backend = match token {
Some(token) => GcsCache::with_token_source(
config,
client,
ResponseCacheCodec,
Arc::new(StaticTokenSource(token)),
),
None => GcsCache::new(config, client, ResponseCacheCodec),
};
let identity = BackendIdentity::Gcs {
bucket_name: backend.bucket_name().to_owned(),
key_prefix: backend.key_prefix().to_owned(),
path_service_account: backend.path_service_account().map(str::to_owned),
};
Self::exact(ResponseCache::new(Arc::new(backend)), identity)
}
pub async fn azure_blob(
account_url: &str,
container: &str,
http: litellm_http::Client,
) -> Result<Self, Error> {
let backend = AzureBlobCache::connect(
account_url,
container,
http,
ResponseCacheCodec,
tokio::runtime::Handle::current(),
)
.await?;
let identity = BackendIdentity::AzureBlob {
account_url: backend.account_url().to_owned(),
container: backend.container_name().to_owned(),
};
Ok(Self::exact(ResponseCache::new(Arc::new(backend)), identity))
}
/// Wraps a built exact backend; the TTL a facade must match comes from the built cache.
fn exact<B>(cache: ResponseCache<B>, identity: BackendIdentity) -> Self
where
ResponseCache<B>: ExactResponseCache + 'static,
B: litellm_cache::BaseCache<Value = litellm_cache_response::CacheEntry>,
B::Context: Default + PartialEq,
{
Self::exact_service(Arc::new(cache), None, identity)
}
/// Wraps an exact backend whose Python class defines `test_connection`.
fn exact_probed<B>(cache: ResponseCache<B>, identity: BackendIdentity) -> Self
where
ResponseCache<B>: ExactResponseCache + ConnectionProbe + 'static,
B: litellm_cache::BaseCache<Value = litellm_cache_response::CacheEntry>,
B::Context: Default + PartialEq,
{
let cache = Arc::new(cache);
Self::exact_service(cache.clone(), Some(cache), identity)
}
fn exact_service(
cache: Arc<dyn ExactResponseCache>,
probe: Option<Arc<dyn ConnectionProbe>>,
identity: BackendIdentity,
) -> Self {
let default_ttl = cache.default_ttl();
let identity = match identity {
BackendIdentity::Memory {
capacity,
max_entry_bytes,
..
} => BackendIdentity::Memory {
capacity,
max_entry_bytes,
default_ttl,
},
BackendIdentity::Redis {
topology,
namespace,
..
} => BackendIdentity::Redis {
topology,
namespace,
default_ttl,
},
other => other,
};
Self::Exact(Arc::new(ExactService {
cache,
probe,
buffer: None,
identity,
}))
}
pub fn valkey_semantic(
url: &str,
similarity_threshold: f64,
index_name: String,
embedder: PythonEmbedder,
) -> Result<Self, Error> {
let backend = ValkeySemanticCache::new(
url,
embedder.clone(),
ResponseCacheCodec,
ValkeySemanticConfig {
similarity_threshold,
index_name,
},
)?;
Ok(Self::ValkeySemantic {
cache: Arc::new(ResponseCache::new(Arc::new(backend))),
embedder,
scope: String::from("key"),
})
}
pub fn redis_semantic(
url: &str,
embedder: PythonEmbedder,
config: RedisSemanticConfig,
) -> Result<Self, Error> {
let backend = RedisSemanticCache::new(url, embedder.clone(), ResponseCacheCodec, config)?;
Ok(Self::RedisSemantic {
cache: Arc::new(ResponseCache::new(Arc::new(backend))),
embedder,
})
}
pub(super) async fn qdrant_semantic(
config: QdrantSemanticCacheConfig,
client: litellm_http::Client,
runtime: tokio::runtime::Handle,
) -> Result<Self, Error> {
let qdrant = qdrant_client::Qdrant::from_url(&config.grpc_url)
.skip_compatibility_check()
.api_key(config.api_key.as_deref())
.build()
.map_err(|_| Error::Unavailable)?;
let qdrant_config = config.to_qdrant_config();
let embedder = OpenAiEmbedder::new(client, config.embedding);
let cache = QdrantSemanticCache::connect(
qdrant,
embedder,
ResponseCacheCodec,
qdrant_config,
runtime,
)
.await?;
Ok(Self::QdrantSemantic(Arc::new(ResponseCache::new(
Arc::new(cache),
))))
}
pub fn identity(&self) -> BackendIdentity {
match self {
Self::Exact(service) => service.identity.clone(),
Self::ValkeySemantic { cache, .. } => BackendIdentity::ValkeySemantic {
index_name: cache.backend().index_name().to_owned(),
similarity_threshold: cache.backend().similarity_threshold(),
},
Self::RedisSemantic { cache, .. } => BackendIdentity::RedisSemantic {
index_name: cache.backend().index_name().to_owned(),
similarity_threshold: cache.backend().similarity_threshold(),
},
Self::QdrantSemantic(cache) => BackendIdentity::QdrantSemantic {
collection_name: cache.backend().collection_name().to_owned(),
similarity_threshold: cache.backend().similarity_threshold(),
vector_size: cache.backend().vector_size(),
embedding_model: cache.backend().embedder().model().to_owned(),
},
}
}
pub fn with_redis_flush_size(self, flush_size: Option<usize>) -> Self {
match self {
Self::Exact(service) if matches!(service.identity, BackendIdentity::Redis { .. }) => {
Self::Exact(Arc::new(ExactService {
cache: Arc::clone(&service.cache),
probe: service.probe.clone(),
buffer: flush_size.map(WriteBuffer::new),
identity: service.identity.clone(),
}))
}
value => value,
}
}
pub fn with_scope(self, scope: String) -> Self {
match self {
Self::ValkeySemantic {
cache, embedder, ..
} => Self::ValkeySemantic {
cache,
embedder,
scope,
},
value => value,
}
}
pub fn embedder_object(&self) -> Option<&Py<PyAny>> {
match self {
Self::RedisSemantic { embedder, .. } => Some(embedder.object()),
_ => None,
}
}
/// The prompt and metadata this backend would embed for `request`, if it has a prompt.
pub(in crate::cache) fn embedding_input(
&self,
request: &NativeRequest,
) -> Option<EmbeddingInput> {
let context = match self {
Self::ValkeySemantic { scope, .. } => request.scoped_semantic(scope).context,
Self::RedisSemantic { .. } => request.semantic().context,
Self::Exact(_) | Self::QdrantSemantic(_) => return None,
};
let prompt = litellm_cache::semantic::prompt_from_context(&context)?;
Some(EmbeddingInput {
prompt,
metadata: context.metadata,
})
}
/// Drives a semantic operation whose embedding comes from Python.
fn python_semantic<'py>(
&self,
py: Python<'py>,
operation: SemanticOperation,
) -> PyResult<Bound<'py, PyAny>> {
let (embedder, failure) = match self {
Self::ValkeySemantic { embedder, .. } => (embedder, EmbeddingFailure::Propagate),
Self::RedisSemantic { embedder, .. } => (embedder, EmbeddingFailure::Unavailable),
Self::Exact(_) | Self::QdrantSemantic(_) => {
return Err(pyo3::exceptions::PyRuntimeError::new_err(
"semantic execution requires a Python-embedded backend",
));
}
};
drive(
py,
SemanticExecution::new(self.clone(), embedder.clone(), failure, operation),
)
}
pub fn lookup(&self, request: &NativeRequest, now: Duration) -> Result<Option<Value>, Error> {
match self {
Self::Exact(service) => service.cache.lookup(&request.exact(), now),
Self::ValkeySemantic { cache, scope, .. } => {
cache.lookup(&request.scoped_semantic(scope), now)
}
Self::RedisSemantic { cache, .. } => cache.lookup(&request.semantic(), now),
Self::QdrantSemantic(cache) => cache.lookup(&request.semantic(), now),
}
}
/// `lookup` plus the similarity Python's semantic backend writes to the request metadata.
/// Exact backends report none.
pub fn lookup_semantic(
&self,
request: &NativeRequest,
now: Duration,
) -> Result<SemanticLookup<Value>, Error> {
match self {
Self::Exact(service) => service
.cache
.lookup(&request.exact(), now)
.map(exact_lookup),
Self::ValkeySemantic { cache, scope, .. } => {
redis_family(cache.lookup_semantic(&request.scoped_semantic(scope), now))
}
Self::RedisSemantic { cache, .. } => {
redis_family(cache.lookup_semantic(&request.semantic(), now))
}
Self::QdrantSemantic(cache) => cache.lookup_semantic(&request.semantic(), now),
}
}
pub fn store(
&self,
request: &NativeRequest,
response: Value,
now: Duration,
) -> Result<(), Error> {
match self {
Self::Exact(service) => service.cache.store(&request.exact(), response, now),
Self::ValkeySemantic { cache, scope, .. } => {
cache.store(&request.scoped_semantic(scope), response, now)
}
Self::RedisSemantic { cache, .. } => cache.store(&request.semantic(), response, now),
Self::QdrantSemantic(cache) => cache.store(&request.semantic(), response, now),
}
}
pub fn lookup_batch(
&self,
requests: &[NativeRequest],
now: Duration,
) -> Result<PartialHits, Error> {
match self {
Self::Exact(service) => service.cache.lookup_batch(&exact_requests(requests), now),
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => {
Err(Error::UnsupportedOperation)
}
}
}
pub async fn async_lookup(
&self,
request: &NativeRequest,
now: Duration,
) -> Result<Option<Value>, Error> {
match self {
Self::Exact(service) => service.cache.async_lookup(&request.exact(), now).await,
Self::ValkeySemantic { cache, scope, .. } => {
cache
.async_lookup(&request.scoped_semantic(scope), now)
.await
}
Self::RedisSemantic { cache, .. } => cache.async_lookup(&request.semantic(), now).await,
Self::QdrantSemantic(cache) => cache.async_lookup(&request.semantic(), now).await,
}
}
pub async fn async_lookup_semantic(
&self,
request: &NativeRequest,
now: Duration,
) -> Result<SemanticLookup<Value>, Error> {
match self {
Self::Exact(service) => service
.cache
.async_lookup(&request.exact(), now)
.await
.map(exact_lookup),
Self::ValkeySemantic { cache, scope, .. } => redis_family(
cache
.async_lookup_semantic(&request.scoped_semantic(scope), now)
.await,
),
Self::RedisSemantic { cache, .. } => {
redis_family(cache.async_lookup_semantic(&request.semantic(), now).await)
}
Self::QdrantSemantic(cache) => {
cache.async_lookup_semantic(&request.semantic(), now).await
}
}
}
pub(in crate::cache) fn async_lookup_semantic_py<'py>(
&self,
py: Python<'py>,
request: NativeRequest,
) -> PyResult<Bound<'py, PyAny>> {
match self {
Self::Exact(_) | Self::QdrantSemantic(_) => {
let service = self.clone();
crate::execution::run_async(
py,
async move {
service
.async_lookup_semantic(&request, now())
.await
.map(SemanticReply::from)
},
cache_error,
)
}
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
self.python_semantic(py, SemanticOperation::LookupSemantic(request))
}
}
}
pub(in crate::cache) fn async_lookup_py<'py>(
&self,
py: Python<'py>,
request: NativeRequest,
) -> PyResult<Bound<'py, PyAny>> {
match self {
Self::Exact(_) | Self::QdrantSemantic(_) => {
let service = self.clone();
crate::execution::run_async(
py,
async move { service.async_lookup(&request, now()).await },
cache_error,
)
}
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
self.python_semantic(py, SemanticOperation::Lookup(request))
}
}
}
pub async fn async_store(
&self,
request: &NativeRequest,
response: Value,
now: Duration,
) -> Result<(), Error> {
match self {
Self::Exact(service) => match &service.buffer {
None => {
service
.cache
.async_store(&request.exact(), response, now)
.await
}
Some(buffer) => {
buffer
.async_store(service.cache.as_ref(), &request.exact(), response, now)
.await
}
},
Self::ValkeySemantic { cache, scope, .. } => {
cache
.async_store(&request.scoped_semantic(scope), response, now)
.await
}
Self::RedisSemantic { cache, .. } => {
cache.async_store(&request.semantic(), response, now).await
}
Self::QdrantSemantic(cache) => {
cache.async_store(&request.semantic(), response, now).await
}
}
}
pub(in crate::cache) fn async_store_py<'py>(
&self,
py: Python<'py>,
request: NativeRequest,
response: Value,
) -> PyResult<Bound<'py, PyAny>> {
match self {
Self::Exact(_) | Self::QdrantSemantic(_) => {
let service = self.clone();
crate::execution::run_async(
py,
async move { service.async_store(&request, response, now()).await },
cache_error,
)
}
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
self.python_semantic(py, SemanticOperation::Store(request, response))
}
}
}
pub async fn async_lookup_batch(
&self,
requests: &[NativeRequest],
now: Duration,
) -> Result<PartialHits, Error> {
match self {
Self::Exact(service) => {
service
.cache
.async_lookup_batch(&exact_requests(requests), now)
.await
}
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => {
Err(Error::UnsupportedOperation)
}
}
}
pub async fn async_store_batch(
&self,
entries: Vec<(NativeRequest, Value)>,
now: Duration,
) -> Result<(), Error> {
match self {
Self::Exact(service) => {
let entries = entries
.into_iter()
.map(|(request, value)| (request.exact(), value))
.collect();
service.cache.async_store_batch(entries, now).await
}
Self::ValkeySemantic { cache, scope, .. } => {
let entries = entries
.into_iter()
.map(|(request, value)| (request.scoped_semantic(scope), value))
.collect();
cache.async_store_batch(entries, now).await
}
Self::RedisSemantic { .. } => Err(Error::UnsupportedOperation),
Self::QdrantSemantic(cache) => {
let entries = entries
.into_iter()
.map(|(request, value)| (request.semantic(), value))
.collect();
cache.async_store_batch(entries, now).await
}
}
}
pub(in crate::cache) fn async_store_batch_py<'py>(
&self,
py: Python<'py>,
entries: Vec<(NativeRequest, Value)>,
) -> PyResult<Bound<'py, PyAny>> {
match self {
Self::Exact(_) | Self::QdrantSemantic(_) => {
let service = self.clone();
crate::execution::run_async(
py,
async move { service.async_store_batch(entries, now()).await },
cache_error,
)
}
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
self.python_semantic(py, SemanticOperation::StoreBatch(entries.into()))
}
}
}
pub async fn async_flush(&self) -> Result<(), Error> {
match self {
Self::Exact(service) => {
if let Some(buffer) = &service.buffer {
buffer.clear()?;
}
service.cache.async_flush().await
}
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => {
Err(Error::UnsupportedOperation)
}
}
}
pub async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
match self {
Self::Exact(service) => match &service.probe {
Some(probe) => probe.test_connection().await,
None => Err(Error::UnsupportedOperation),
},
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => {
Err(Error::UnsupportedOperation)
}
}
}
}
fn exact_requests(requests: &[NativeRequest]) -> Vec<litellm_cache_response::ResponseCacheRequest> {
requests.iter().map(NativeRequest::exact).collect()
}
/// What `lookup_semantic` hands Python: the response and the similarity to stamp, if any.
#[derive(serde::Serialize)]
pub(in crate::cache) struct SemanticReply(
pub(in crate::cache) Option<Value>,
pub(in crate::cache) Option<f64>,
);
impl From<SemanticLookup<Value>> for SemanticReply {
fn from(lookup: SemanticLookup<Value>) -> Self {
Self(lookup.value, lookup.similarity)
}
}
fn exact_lookup(value: Option<Value>) -> SemanticLookup<Value> {
SemanticLookup {
value,
similarity: None,
}
}
/// Python's Redis and Valkey semantic caches catch every lookup failure and stamp `0.0`.
fn redis_family(
lookup: Result<SemanticLookup<Value>, Error>,
) -> Result<SemanticLookup<Value>, Error> {
Ok(lookup.unwrap_or_else(|_| SemanticLookup::miss(Some(0.0))))
}

File diff suppressed because it is too large Load diff

View file

@ -1,129 +0,0 @@
use std::future::Future;
use litellm_cache::Error;
use litellm_host_python::to_py;
use pyo3::{PyTraverseError, PyVisit, prelude::*, types::PyDict};
use serde_json::Value;
tokio::task_local! {
static PREPARED_EMBEDDING: Result<Vec<f32>, Error>;
}
/// Runs `future` with the vector the Python embedder already produced, so the backend's
/// `async_embed` never has to call back into Python from the runtime.
pub(in crate::cache) fn with_prepared_embedding<F: Future>(
vector: Result<Vec<f32>, Error>,
future: F,
) -> impl Future<Output = F::Output> {
PREPARED_EMBEDDING.scope(vector, future)
}
/// The Python object that owns embedding for a semantic backend.
pub(in crate::cache) struct PythonEmbedder(Py<PyAny>);
impl Clone for PythonEmbedder {
fn clone(&self) -> Self {
Python::attach(|py| Self(self.0.clone_ref(py)))
}
}
impl PythonEmbedder {
pub(in crate::cache) fn new(object: Py<PyAny>) -> Self {
Self(object)
}
pub(in crate::cache) fn object(&self) -> &Py<PyAny> {
&self.0
}
pub(in crate::cache) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.0)
}
fn metadata_kwargs<'py>(
py: Python<'py>,
metadata: Option<&Value>,
) -> PyResult<Bound<'py, PyDict>> {
let kwargs = PyDict::new(py);
kwargs.set_item("metadata", to_py(py, &metadata)?)?;
Ok(kwargs)
}
/// The awaitable of `_get_async_embedding(prompt, metadata=...)`, to run in the caller's loop.
pub(in crate::cache) fn async_embedding(
&self,
py: Python<'_>,
prompt: &str,
metadata: Option<&Value>,
) -> PyResult<Py<PyAny>> {
let kwargs = Self::metadata_kwargs(py, metadata)?;
self.0
.bind(py)
.call_method("_get_async_embedding", (prompt,), Some(&kwargs))
.map(Bound::unbind)
}
pub(in crate::cache) fn extract(vector: Bound<'_, PyAny>) -> PyResult<Vec<f32>> {
Ok(vector
.extract::<Vec<f64>>()?
.into_iter()
.map(|value| value as f32)
.collect())
}
fn embed_sync(&self, prompt: &str, metadata: Option<&Value>) -> Result<Vec<f32>, Error> {
Python::attach(|py| {
let kwargs = Self::metadata_kwargs(py, metadata)?;
Self::extract(self.0.bind(py).call_method(
"_get_embedding",
(prompt,),
Some(&kwargs),
)?)
})
.map_err(|_| Error::Unavailable)
}
fn seeded_embedding() -> Result<Vec<f32>, Error> {
PREPARED_EMBEDDING
.try_with(Clone::clone)
.unwrap_or(Err(Error::Unavailable))
}
}
impl litellm_cache::semantic::Embedder for PythonEmbedder {
fn embed(&self, prompt: &str, metadata: Option<&Value>) -> Result<Vec<f32>, Error> {
self.embed_sync(prompt, metadata)
}
fn async_embed(
&self,
_prompt: &str,
_metadata: Option<&Value>,
) -> impl Future<Output = Result<Vec<f32>, Error>> + Send {
std::future::ready(Self::seeded_embedding())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn async_embed_returns_the_seeded_vector_or_unavailable() {
Python::initialize();
let object = Python::attach(|py| py.None());
let embedder = PythonEmbedder::new(object);
let scoped_embedder = embedder.clone();
let scoped = with_prepared_embedding(Ok(vec![0.25]), async move {
litellm_cache::semantic::Embedder::async_embed(&scoped_embedder, "prompt", None).await
});
assert_eq!(scoped.await, Ok(vec![0.25]));
let unscoped =
litellm_cache::semantic::Embedder::async_embed(&embedder, "prompt", None).await;
assert_eq!(unscoped, Err(Error::Unavailable));
let valkey = with_prepared_embedding(Ok(vec![0.5]), async move {
litellm_cache::semantic::Embedder::async_embed(&embedder, "prompt", None).await
});
assert_eq!(valkey.await, Ok(vec![0.5]));
}
}

View file

@ -1,502 +0,0 @@
use litellm_cache_redis::RedisTopology;
use litellm_host_python::from_py;
use pyo3::{
PyTraverseError, PyVisit,
exceptions::PyTypeError,
prelude::*,
types::{PyDict, PyTuple, PyType},
};
use serde_json::Value;
use super::{
backend::NativeResponseCache,
config::{CacheConfigProjection, NativeCacheConfig},
identity::BackendIdentity,
};
struct ClassGuard {
class: Py<PyType>,
attributes: Vec<(String, Py<PyAny>)>,
}
struct ObjectGuard {
reference: Py<PyAny>,
classes: Vec<ClassGuard>,
config_names: &'static [&'static str],
config: Vec<Value>,
}
struct RedisPoolGuard {
reference: Py<PyAny>,
connection_class: Py<PyAny>,
connection_kwargs: Py<PyAny>,
max_connections: Option<usize>,
client_name: &'static str,
attributes: RedisPoolAttributes,
}
struct DiskStoreGuard {
reference: Py<PyAny>,
directory: String,
}
struct AzureBlobClientGuard {
sync_client: Py<PyAny>,
async_client: Py<PyAny>,
url: String,
container_name: String,
}
struct S3ClientGuard {
reference: Py<PyAny>,
}
enum ConnectionGuard {
None,
RedisPool(RedisPoolGuard),
AzureBlob(AzureBlobClientGuard),
S3(S3ClientGuard),
}
#[derive(Clone, Copy)]
struct RedisPoolAttributes {
pool: &'static str,
connection_class: &'static str,
max_connections: Option<&'static str>,
}
const STANDALONE_POOL: RedisPoolAttributes = RedisPoolAttributes {
pool: "connection_pool",
connection_class: "connection_class",
max_connections: Some("max_connections"),
};
const CLUSTER_POOL: RedisPoolAttributes = RedisPoolAttributes {
pool: "nodes_manager",
connection_class: "connection_pool_class",
max_connections: None,
};
const VALKEY_POOL: RedisPoolAttributes = STANDALONE_POOL;
/// Class-level defaults an instance overwrites with its own state rather than behavior:
/// `Cache._native_cache` holds the runtime `Cache.__init__` resolved.
const INSTANCE_STATE: &[&str] = &["_native_cache"];
pub(in crate::cache) struct FacadeGuard {
outer: ObjectGuard,
backend: ObjectGuard,
disk_store: Option<DiskStoreGuard>,
connection: ConnectionGuard,
}
impl ObjectGuard {
fn capture(
py: Python<'_>,
object: &Bound<'_, PyAny>,
config_names: &'static [&'static str],
) -> PyResult<Self> {
let classes = object
.get_type()
.getattr("__mro__")?
.cast_into::<PyTuple>()?
.iter()
.map(|class| {
let class = class.cast_into::<PyType>()?;
let attributes = class
.getattr("__dict__")?
.call_method0("items")?
.try_iter()?
.map(|item| item?.extract::<(String, Py<PyAny>)>())
.collect::<PyResult<Vec<_>>>()?;
Ok(ClassGuard {
class: class.unbind(),
attributes,
})
})
.collect::<PyResult<Vec<_>>>()?;
let guard = Self {
reference: py
.import("weakref")?
.getattr("ref")?
.call1((object,))?
.unbind(),
classes,
config_names,
config: Self::config(object, config_names)?,
};
if !guard.matches(py, object)? {
return Err(PyTypeError::new_err(
"native facade registration requires unmodified built-in methods",
));
}
Ok(guard)
}
fn config(object: &Bound<'_, PyAny>, names: &[&str]) -> PyResult<Vec<Value>> {
names
.iter()
.map(|name| match object.getattr(*name) {
Ok(value) => from_py(&value),
Err(error)
if error.is_instance_of::<pyo3::exceptions::PyAttributeError>(object.py()) =>
{
Ok(Value::Null)
}
Err(error) => Err(error),
})
.collect()
}
fn matches(&self, py: Python<'_>, object: &Bound<'_, PyAny>) -> PyResult<bool> {
if !self.reference.bind(py).call0()?.is(object) {
return Ok(false);
}
let mro = object
.get_type()
.getattr("__mro__")?
.cast_into::<PyTuple>()?;
if mro.len() != self.classes.len() {
return Ok(false);
}
let instance = object.getattr("__dict__")?.cast_into::<PyDict>()?;
for (class, expected) in mro.iter().zip(&self.classes) {
if !class.is(expected.class.bind(py)) {
return Ok(false);
}
let attributes = class.getattr("__dict__")?;
if attributes.len()? != expected.attributes.len() {
return Ok(false);
}
for (name, value) in &expected.attributes {
if (instance.contains(name)?
&& !self.config_names.contains(&name.as_str())
&& !INSTANCE_STATE.contains(&name.as_str()))
|| !attributes.get_item(name)?.is(value.bind(py))
{
return Ok(false);
}
}
}
Ok(Self::config(object, self.config_names)? == self.config)
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.reference)?;
for class in &self.classes {
visit.call(&class.class)?;
for (_, value) in &class.attributes {
visit.call(value)?;
}
}
Ok(())
}
}
impl RedisPoolGuard {
fn capture(
backend: &Bound<'_, PyAny>,
client_name: &'static str,
attributes: RedisPoolAttributes,
) -> PyResult<Self> {
let pool = backend.getattr(client_name)?.getattr(attributes.pool)?;
Ok(Self {
reference: pool.clone().unbind(),
connection_class: pool.getattr(attributes.connection_class)?.unbind(),
connection_kwargs: pool
.getattr("connection_kwargs")?
.call_method0("copy")?
.unbind(),
max_connections: attributes
.max_connections
.map(|name| pool.getattr(name)?.extract::<usize>())
.transpose()?,
client_name,
attributes,
})
}
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
let pool = backend
.getattr(self.client_name)?
.getattr(self.attributes.pool)?;
Ok(self.reference.bind(py).is(&pool)
&& self
.connection_class
.bind(py)
.is(&pool.getattr(self.attributes.connection_class)?)
&& self.max_connections
== self
.attributes
.max_connections
.map(|name| pool.getattr(name)?.extract::<usize>())
.transpose()?
&& self
.connection_kwargs
.bind(py)
.eq(pool.getattr("connection_kwargs")?)?)
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.reference)?;
visit.call(&self.connection_class)?;
visit.call(&self.connection_kwargs)
}
}
impl DiskStoreGuard {
fn capture(backend: &Bound<'_, PyAny>) -> PyResult<Self> {
let store = backend.getattr("disk_cache")?;
Ok(Self {
reference: store.clone().unbind(),
directory: store.getattr("directory")?.extract()?,
})
}
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
let store = backend.getattr("disk_cache")?;
Ok(self.reference.bind(py).is(&store)
&& self.directory == store.getattr("directory")?.extract::<String>()?)
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.reference)
}
}
impl AzureBlobClientGuard {
fn capture(backend: &Bound<'_, PyAny>) -> PyResult<Self> {
let sync_client = backend.getattr("container_client")?;
Ok(Self {
url: sync_client.getattr("url")?.extract::<String>()?,
container_name: sync_client.getattr("container_name")?.extract::<String>()?,
sync_client: sync_client.unbind(),
async_client: backend.getattr("async_container_client")?.unbind(),
})
}
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
let sync_client = backend.getattr("container_client")?;
Ok(self.sync_client.bind(py).is(&sync_client)
&& self
.async_client
.bind(py)
.is(&backend.getattr("async_container_client")?)
&& self.url == sync_client.getattr("url")?.extract::<String>()?
&& self.container_name == sync_client.getattr("container_name")?.extract::<String>()?)
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.sync_client)?;
visit.call(&self.async_client)
}
}
impl S3ClientGuard {
fn capture(backend: &Bound<'_, PyAny>) -> PyResult<Self> {
Ok(Self {
reference: backend.getattr("s3_client")?.unbind(),
})
}
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
Ok(self.reference.bind(py).is(&backend.getattr("s3_client")?))
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.reference)
}
}
impl ConnectionGuard {
fn capture(kind: &str, cluster: bool, backend: &Bound<'_, PyAny>) -> PyResult<Self> {
Ok(match (kind, cluster) {
("redis", false) => Self::RedisPool(RedisPoolGuard::capture(
backend,
"redis_client",
STANDALONE_POOL,
)?),
("redis", true) => Self::RedisPool(RedisPoolGuard::capture(
backend,
"redis_client",
CLUSTER_POOL,
)?),
("valkey-semantic", _) => Self::RedisPool(RedisPoolGuard::capture(
backend,
"sync_client",
VALKEY_POOL,
)?),
("disk", _) => Self::None,
("azure-blob", _) => Self::AzureBlob(AzureBlobClientGuard::capture(backend)?),
("s3", _) => Self::S3(S3ClientGuard::capture(backend)?),
_ => Self::None,
})
}
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
match self {
Self::None => Ok(true),
Self::RedisPool(guard) => guard.matches(py, backend),
Self::AzureBlob(guard) => guard.matches(py, backend),
Self::S3(guard) => guard.matches(py, backend),
}
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
match self {
Self::None => Ok(()),
Self::RedisPool(guard) => guard.traverse(visit),
Self::AzureBlob(guard) => guard.traverse(visit),
Self::S3(guard) => guard.traverse(visit),
}
}
}
impl FacadeGuard {
pub(in crate::cache) fn capture(
py: Python<'_>,
facade: &Bound<'_, PyAny>,
service: &NativeResponseCache,
) -> PyResult<Self> {
let identity = service.identity();
let kind = identity.kind();
let cache_type = py.import("litellm.caching.caching")?.getattr("Cache")?;
if !facade.get_type().is(&cache_type) {
return Err(PyTypeError::new_err(
"only exact built-in Cache facades can be registered",
));
}
let cluster = matches!(
identity,
BackendIdentity::Redis {
topology: RedisTopology::Cluster { .. },
..
}
);
let (module, name) = match (kind, cluster) {
("memory", _) => ("litellm.caching.in_memory_cache", "InMemoryCache"),
("redis", false) => ("litellm.caching.redis_cache", "RedisCache"),
("redis", true) => ("litellm.caching.redis_cluster_cache", "RedisClusterCache"),
("redis_semantic", _) => ("litellm.caching.redis_semantic_cache", "RedisSemanticCache"),
("qdrant_semantic", _) => (
"litellm.caching.qdrant_semantic_cache",
"QdrantSemanticCache",
),
("gcs", _) => ("litellm.caching.gcs_cache", "GCSCache"),
("valkey-semantic", _) => (
"litellm.caching.valkey_semantic_cache",
"ValkeySemanticCache",
),
("disk", _) => ("litellm.caching.disk_cache", "DiskCache"),
("azure-blob", _) => ("litellm.caching.azure_blob_cache", "AzureBlobCache"),
("s3", _) => ("litellm.caching.s3_cache", "S3Cache"),
_ => unreachable!(),
};
let cache_kind = identity.cache_type();
let backend = facade.getattr("cache")?;
if facade.getattr("type")?.extract::<String>()? != cache_kind
|| !backend.get_type().is(&py.import(module)?.getattr(name)?)
{
return Err(PyTypeError::new_err(
"facade and native backend types must match",
));
}
let config = match NativeCacheConfig::project(facade)? {
CacheConfigProjection::Native(config) => *config,
CacheConfigProjection::Unsupported(reason) => {
return Err(PyTypeError::new_err(reason.message()));
}
};
if let Some(message) = config.service_mismatch(service) {
return Err(PyTypeError::new_err(message));
}
if kind == "redis_semantic"
&& service
.embedder_object()
.is_none_or(|embedder| !backend.is(embedder.bind(py)))
{
return Err(PyTypeError::new_err(
"facade backend must be the native embedder",
));
}
Ok(Self {
outer: ObjectGuard::capture(
py,
facade,
&[
"type",
"mode",
"ttl",
"namespace",
"supported_call_types",
"redis_flush_size",
"semantic_cache_scope",
],
)?,
backend: ObjectGuard::capture(
py,
&backend,
&[
"namespace",
"default_ttl",
"max_size_in_memory",
"max_size_per_item",
"redis_kwargs",
"redis_flush_size",
"similarity_threshold",
"distance_threshold",
"embedding_model",
"embedding_max_input_tokens",
"embedding_timeout",
"qdrant_api_base",
"qdrant_api_key",
"collection_name",
"vector_size",
"_index_name",
"_redis_url",
"similarity_threshold",
"embedding_model",
"index_name",
"embedding_max_input_tokens",
"embedding_timeout",
"bucket_name",
"key_prefix",
"path_service_account",
],
)?,
disk_store: (kind == "disk")
.then(|| DiskStoreGuard::capture(&backend))
.transpose()?,
connection: ConnectionGuard::capture(kind, cluster, &backend)?,
})
}
pub(in crate::cache) fn matches(
&self,
py: Python<'_>,
facade: &Bound<'_, PyAny>,
) -> PyResult<bool> {
if !self.outer.matches(py, facade)? {
return Ok(false);
}
let backend = facade.getattr("cache")?;
if !self.backend.matches(py, &backend)? {
return Ok(false);
}
if let Some(guard) = &self.disk_store
&& !guard.matches(py, &backend)?
{
return Ok(false);
}
self.connection.matches(py, &backend)
}
pub(in crate::cache) fn traverse(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
self.outer.traverse(&visit)?;
self.backend.traverse(&visit)?;
if let Some(guard) = &self.disk_store {
guard.traverse(&visit)?;
}
self.connection.traverse(&visit)
}
}

View file

@ -1,510 +0,0 @@
use std::{path::PathBuf, time::Duration};
use litellm_cache_redis::RedisTopology;
/// What makes a native backend the one a Python facade describes: the configuration a user can
/// observe on the Python object, captured once so facade projection and native construction
/// compare plain data instead of reaching into each backend type.
#[derive(Clone, Debug, PartialEq)]
pub(in crate::cache) enum BackendIdentity {
Memory {
capacity: usize,
max_entry_bytes: Option<usize>,
default_ttl: Option<Duration>,
},
Redis {
topology: RedisTopology,
namespace: Option<String>,
default_ttl: Option<Duration>,
},
S3 {
bucket: String,
key_prefix: String,
region: String,
endpoint: Option<String>,
},
Gcs {
bucket_name: String,
key_prefix: String,
path_service_account: Option<String>,
},
Disk {
directory: PathBuf,
},
AzureBlob {
account_url: String,
container: String,
},
RedisSemantic {
index_name: String,
/// The backend stores the threshold as `f32`; a facade's `f64` is compared at that width.
similarity_threshold: f32,
},
ValkeySemantic {
index_name: String,
similarity_threshold: f64,
},
QdrantSemantic {
collection_name: String,
similarity_threshold: f64,
vector_size: u64,
embedding_model: String,
},
}
const TYPES: &str = "facade and native backend types must match";
impl BackendIdentity {
pub(in crate::cache) fn kind(&self) -> &'static str {
match self {
Self::Memory { .. } => "memory",
Self::Redis { .. } => "redis",
Self::S3 { .. } => "s3",
Self::Gcs { .. } => "gcs",
Self::ValkeySemantic { .. } => "valkey-semantic",
Self::RedisSemantic { .. } => "redis_semantic",
Self::QdrantSemantic { .. } => "qdrant_semantic",
Self::Disk { .. } => "disk",
Self::AzureBlob { .. } => "azure-blob",
}
}
/// The `LiteLLMCacheType` value a facade of this backend carries in `Cache.type`.
pub(in crate::cache) fn cache_type(&self) -> &'static str {
match self {
Self::Memory { .. } => "local",
Self::Redis { .. } => "redis",
Self::S3 { .. } => "s3",
Self::Gcs { .. } => "gcs",
Self::ValkeySemantic { .. } => "valkey-semantic",
Self::RedisSemantic { .. } => "redis-semantic",
Self::QdrantSemantic { .. } => "qdrant-semantic",
Self::Disk { .. } => "disk",
Self::AzureBlob { .. } => "azure-blob",
}
}
/// The first difference between the facade's configuration (`self`) and the native
/// backend (`native`), in the order Python users see the attributes.
pub(in crate::cache) fn mismatch(&self, native: &Self) -> Option<&'static str> {
let mut differences: Vec<(bool, &'static str)> = Vec::new();
let mut differs = |condition: bool, message: &'static str| {
differences.push((condition, message));
};
match (self, native) {
(
Self::Memory {
capacity,
max_entry_bytes,
default_ttl,
},
Self::Memory {
capacity: native_capacity,
max_entry_bytes: native_max_entry_bytes,
default_ttl: native_default_ttl,
},
) => {
differs(
default_ttl != native_default_ttl,
"facade and native backend default TTLs must match",
);
differs(
capacity != native_capacity,
"facade and native backend capacities must match",
);
differs(
max_entry_bytes != native_max_entry_bytes,
"facade and native backend item limits must match",
);
}
(
Self::Redis {
topology,
namespace,
default_ttl,
},
Self::Redis {
topology: native_topology,
namespace: native_namespace,
default_ttl: native_default_ttl,
},
) => {
differs(
default_ttl != native_default_ttl,
"facade and native backend default TTLs must match",
);
differs(
topology != native_topology,
"facade and native backend topologies must match",
);
differs(
namespace != native_namespace,
"facade and native backend namespaces must match",
);
}
(
Self::S3 {
bucket,
key_prefix,
region,
endpoint,
},
Self::S3 {
bucket: native_bucket,
key_prefix: native_key_prefix,
region: native_region,
endpoint: native_endpoint,
},
) => {
differs(
bucket != native_bucket,
"facade and native backend buckets must match",
);
differs(
key_prefix != native_key_prefix,
"facade and native backend key prefixes must match",
);
differs(
region != native_region,
"facade and native backend regions must match",
);
differs(
endpoint != native_endpoint,
"facade and native backend endpoints must match",
);
}
(
Self::Gcs {
bucket_name,
key_prefix,
path_service_account,
},
Self::Gcs {
bucket_name: native_bucket_name,
key_prefix: native_key_prefix,
path_service_account: native_path_service_account,
},
) => {
differs(
bucket_name != native_bucket_name,
"facade and native backend buckets must match",
);
differs(
key_prefix != native_key_prefix,
"facade and native backend key prefixes must match",
);
differs(
path_service_account != native_path_service_account,
"facade and native backend credentials must match",
);
}
(
Self::Disk { directory },
Self::Disk {
directory: native_directory,
},
) => {
let canonical = |path: &PathBuf| std::fs::canonicalize(path).ok();
differs(
canonical(directory) != canonical(native_directory),
"facade and native backend directories must match",
);
}
(
Self::AzureBlob {
account_url,
container,
},
Self::AzureBlob {
account_url: native_account_url,
container: native_container,
},
) => {
differs(
account_url != native_account_url || container != native_container,
"facade and native backend containers must match",
);
}
(
Self::RedisSemantic {
index_name,
similarity_threshold,
},
Self::RedisSemantic {
index_name: native_index_name,
similarity_threshold: native_similarity_threshold,
},
) => {
differs(
index_name != native_index_name,
"facade and native backend index names must match",
);
differs(
similarity_threshold != native_similarity_threshold,
"facade and native backend similarity thresholds must match",
);
}
(
Self::ValkeySemantic {
index_name,
similarity_threshold,
},
Self::ValkeySemantic {
index_name: native_index_name,
similarity_threshold: native_similarity_threshold,
},
) => {
differs(
index_name != native_index_name
|| similarity_threshold != native_similarity_threshold,
"facade and native semantic settings must match",
);
}
(
Self::QdrantSemantic {
collection_name,
similarity_threshold,
vector_size,
embedding_model,
},
Self::QdrantSemantic {
collection_name: native_collection_name,
similarity_threshold: native_similarity_threshold,
vector_size: native_vector_size,
embedding_model: native_embedding_model,
},
) => {
differs(
collection_name != native_collection_name,
"facade and native backend collections must match",
);
differs(
similarity_threshold != native_similarity_threshold,
"facade and native backend similarity thresholds must match",
);
differs(
vector_size != native_vector_size,
"facade and native backend vector sizes must match",
);
differs(
embedding_model != native_embedding_model,
"facade and native backend embedding models must match",
);
}
_ => return Some(TYPES),
}
differences
.into_iter()
.find_map(|(condition, message)| condition.then_some(message))
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use litellm_cache_redis::{RedisNode, RedisTopology};
use super::BackendIdentity;
fn memory() -> BackendIdentity {
BackendIdentity::Memory {
capacity: 200,
max_entry_bytes: Some(1024),
default_ttl: Some(Duration::from_secs(60)),
}
}
fn redis() -> BackendIdentity {
BackendIdentity::Redis {
topology: RedisTopology::Standalone,
namespace: Some("team".into()),
default_ttl: Some(Duration::from_secs(60)),
}
}
fn s3() -> BackendIdentity {
BackendIdentity::S3 {
bucket: "bucket".into(),
key_prefix: "cache/".into(),
region: "us-east-1".into(),
endpoint: None,
}
}
fn gcs() -> BackendIdentity {
BackendIdentity::Gcs {
bucket_name: "bucket".into(),
key_prefix: "cache/".into(),
path_service_account: Some("credentials.json".into()),
}
}
fn azure() -> BackendIdentity {
BackendIdentity::AzureBlob {
account_url: "https://account.blob.core.windows.net".into(),
container: "cache".into(),
}
}
fn redis_semantic() -> BackendIdentity {
BackendIdentity::RedisSemantic {
index_name: "idx".into(),
similarity_threshold: 0.8,
}
}
#[test]
fn redis_semantic_thresholds_compare_at_backend_precision() {
let facade = BackendIdentity::RedisSemantic {
index_name: "idx".into(),
similarity_threshold: 0.8_f64 as f32,
};
assert_eq!(facade.mismatch(&redis_semantic()), None);
}
fn valkey_semantic() -> BackendIdentity {
BackendIdentity::ValkeySemantic {
index_name: "idx".into(),
similarity_threshold: 0.8,
}
}
fn qdrant() -> BackendIdentity {
BackendIdentity::QdrantSemantic {
collection_name: "collection".into(),
similarity_threshold: 0.8,
vector_size: 1536,
embedding_model: "text-embedding-3-small".into(),
}
}
#[test]
fn identical_identities_have_no_mismatch() {
for identity in [
memory(),
redis(),
s3(),
gcs(),
azure(),
redis_semantic(),
valkey_semantic(),
qdrant(),
BackendIdentity::Disk {
directory: std::env::temp_dir(),
},
] {
assert_eq!(identity.mismatch(&identity), None, "{identity:?}");
}
}
#[test]
fn different_kinds_report_a_type_mismatch() {
assert_eq!(
memory().mismatch(&redis()),
Some("facade and native backend types must match")
);
assert_eq!(
redis_semantic().mismatch(&valkey_semantic()),
Some("facade and native backend types must match")
);
}
#[test]
fn the_first_differing_field_names_the_mismatch() {
let BackendIdentity::Memory { capacity, .. } = memory() else {
unreachable!()
};
assert_eq!(
memory().mismatch(&BackendIdentity::Memory {
capacity: capacity + 1,
max_entry_bytes: Some(1),
default_ttl: Some(Duration::from_secs(60)),
}),
Some("facade and native backend capacities must match")
);
assert_eq!(
memory().mismatch(&BackendIdentity::Memory {
capacity,
max_entry_bytes: Some(1),
default_ttl: Some(Duration::from_secs(61)),
}),
Some("facade and native backend default TTLs must match")
);
assert_eq!(
redis().mismatch(&BackendIdentity::Redis {
topology: RedisTopology::Cluster {
startup_nodes: vec![RedisNode {
host: "node".into(),
port: 7000,
}],
},
namespace: None,
default_ttl: Some(Duration::from_secs(60)),
}),
Some("facade and native backend topologies must match")
);
assert_eq!(
s3().mismatch(&BackendIdentity::S3 {
bucket: "bucket".into(),
key_prefix: "cache/".into(),
region: "us-east-1".into(),
endpoint: Some("http://localhost:9000".into()),
}),
Some("facade and native backend endpoints must match")
);
assert_eq!(
gcs().mismatch(&BackendIdentity::Gcs {
bucket_name: "bucket".into(),
key_prefix: "cache/".into(),
path_service_account: None,
}),
Some("facade and native backend credentials must match")
);
assert_eq!(
azure().mismatch(&BackendIdentity::AzureBlob {
account_url: "https://account.blob.core.windows.net".into(),
container: "other".into(),
}),
Some("facade and native backend containers must match")
);
assert_eq!(
valkey_semantic().mismatch(&BackendIdentity::ValkeySemantic {
index_name: "idx".into(),
similarity_threshold: 0.9,
}),
Some("facade and native semantic settings must match")
);
assert_eq!(
qdrant().mismatch(&BackendIdentity::QdrantSemantic {
collection_name: "collection".into(),
similarity_threshold: 0.8,
vector_size: 1536,
embedding_model: "text-embedding-3-large".into(),
}),
Some("facade and native backend embedding models must match")
);
}
#[test]
fn disk_directories_compare_canonically() {
let directory = std::env::temp_dir();
let mut indirect = directory.clone();
indirect.push(".");
assert_eq!(
BackendIdentity::Disk {
directory: directory.clone()
}
.mismatch(&BackendIdentity::Disk {
directory: indirect
}),
None
);
assert_eq!(
BackendIdentity::Disk { directory }.mismatch(&BackendIdentity::Disk {
directory: "/definitely/missing".into()
}),
Some("facade and native backend directories must match")
);
}
}

View file

@ -1,11 +1,3 @@
pub(super) mod activation;
pub(super) mod backend;
pub(super) mod config;
mod embedder;
pub(super) mod facade;
mod identity;
pub(super) mod request;
mod semantic;
pub(super) mod v2;
pub(crate) use v2::NativeCacheHandle;

View file

@ -1,258 +0,0 @@
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use litellm_cache::{ExactCacheContext, SemanticCacheContext};
use litellm_cache_response::{CacheControls, CacheKeyField, CacheKeyInput, ResponseCacheRequest};
use litellm_host_python::from_py;
use pyo3::{exceptions::PyValueError, prelude::*};
use serde::Deserialize;
use serde_json::Value;
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct RequestInput {
key: CacheKeyInput,
controls: Option<CacheControls>,
ttl_seconds: Option<f64>,
max_age_seconds: Option<f64>,
messages: Option<Value>,
input: Option<Value>,
metadata: Option<Value>,
litellm_metadata: Option<Value>,
litellm_params: Option<Value>,
scope: Option<String>,
}
#[derive(Clone)]
pub(in crate::cache) struct NativeRequest {
pub(super) key: CacheKeyInput,
pub(super) controls: CacheControls,
pub(super) ttl: Option<Duration>,
pub(super) max_age: Option<Duration>,
pub(super) messages: Option<Value>,
pub(super) input: Option<Value>,
pub(super) metadata: Option<Value>,
pub(super) litellm_metadata: Option<Value>,
pub(super) litellm_params: Option<Value>,
pub(super) scope: Option<String>,
}
impl NativeRequest {
pub(super) fn exact(&self) -> ResponseCacheRequest<ExactCacheContext> {
ResponseCacheRequest {
key: self.key.clone(),
controls: self.controls,
context: ExactCacheContext { ttl: self.ttl },
max_age: self.max_age,
}
}
/// The request as a semantic backend that keys on the caller's scope sees it.
pub(super) fn semantic(&self) -> ResponseCacheRequest<SemanticCacheContext> {
self.semantic_with(self.key.clone(), self.scope.clone())
}
/// The request keyed the way Python's Valkey semantic cache keys it: prompt fields drop out
/// and the tenant identifiers for `scope` join the key.
pub(super) fn scoped_semantic(
&self,
scope: &str,
) -> ResponseCacheRequest<SemanticCacheContext> {
self.semantic_with(semantic_key(self, scope), Some(scope.to_owned()))
}
fn semantic_with(
&self,
key: CacheKeyInput,
scope: Option<String>,
) -> ResponseCacheRequest<SemanticCacheContext> {
ResponseCacheRequest {
key,
controls: self.controls,
context: SemanticCacheContext {
input: self.input.clone(),
messages: self.messages.clone(),
metadata: self.metadata.clone(),
scope,
ttl: self.ttl,
},
max_age: self.max_age,
}
}
}
fn semantic_key(request: &NativeRequest, scope: &str) -> CacheKeyInput {
let mut key = request.key.clone();
if key.preset.is_some() {
return key;
}
key.fields
.retain(|field| !matches!(field.name.as_str(), "messages" | "prompt" | "input"));
const TENANT: [&str; 3] = [
"user_api_key",
"user_api_key_team_id",
"user_api_key_org_id",
];
let end_user = (scope == "end_user").then_some("user_api_key_end_user_id");
for name in TENANT.into_iter().chain(end_user) {
let sources = [
request.metadata.as_ref(),
request.litellm_metadata.as_ref(),
request
.litellm_params
.as_ref()
.and_then(|params| params.get("metadata")),
request
.litellm_params
.as_ref()
.and_then(|params| params.get("litellm_metadata")),
];
let Some(value) = sources.into_iter().flatten().find_map(|source| {
source
.as_object()
.and_then(|values| values.get(name))
.filter(|value| !value.is_null())
}) else {
continue;
};
let value = match value {
Value::Null => continue,
Value::String(text) => text.clone(),
other => other.to_string(),
};
key.fields.push(CacheKeyField {
name: name.to_owned(),
value: Some(value),
api_parameter: true,
internal_parameter: false,
});
}
key
}
pub(in crate::cache) fn request(value: &Bound<'_, PyAny>) -> PyResult<NativeRequest> {
let input: RequestInput = from_py(value)?;
request_input(input)
}
fn request_input(input: RequestInput) -> PyResult<NativeRequest> {
let controls = input.controls.unwrap_or_else(|| {
ResponseCacheRequest::<ExactCacheContext>::new(input.key.clone()).controls
});
Ok(NativeRequest {
key: input.key,
controls,
ttl: input.ttl_seconds.map(duration).transpose()?,
max_age: input.max_age_seconds.map(duration).transpose()?,
messages: input.messages,
input: input.input,
metadata: input.metadata,
litellm_metadata: input.litellm_metadata,
litellm_params: input.litellm_params,
scope: input.scope,
})
}
pub(in crate::cache) fn requests(value: &Bound<'_, PyAny>) -> PyResult<Vec<NativeRequest>> {
from_py::<Vec<RequestInput>>(value)?
.into_iter()
.map(request_input)
.collect()
}
pub(super) fn duration(seconds: f64) -> PyResult<Duration> {
Duration::try_from_secs_f64(seconds)
.map_err(|_| PyValueError::new_err("cache durations must be finite and nonnegative"))
}
pub(in crate::cache) fn now() -> Duration {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
}
#[cfg(test)]
mod tests {
use litellm_cache_response::{CacheControls, CacheKeyInput, cache_key};
use serde_json::json;
use sha2::{Digest, Sha256};
use super::*;
fn native_request(key: CacheKeyInput, metadata: Value) -> NativeRequest {
NativeRequest {
key,
controls: CacheControls::default(),
ttl: None,
max_age: None,
messages: Some(json!([{"role": "user", "content": "prompt"}])),
input: None,
metadata: Some(metadata),
litellm_metadata: None,
litellm_params: None,
scope: None,
}
}
#[test]
fn semantic_key_matches_python_scope_material() {
let key = CacheKeyInput {
fields: vec![
CacheKeyField {
name: "model".to_owned(),
value: Some("gpt-4.1".to_owned()),
api_parameter: true,
internal_parameter: false,
},
CacheKeyField {
name: "messages".to_owned(),
value: Some("prompt".to_owned()),
api_parameter: true,
internal_parameter: false,
},
],
..Default::default()
};
let request = native_request(
key,
json!({"user_api_key": "k1", "user_api_key_team_id": null}),
);
let expected = format!("{:x}", Sha256::digest(b"model: gpt-4.1user_api_key: k1"));
assert_eq!(cache_key(&semantic_key(&request, "key")), expected);
assert_eq!(cache_key(&request.scoped_semantic("key").key), expected);
let end_user_request = native_request(
request.key.clone(),
json!({"user_api_key": "k1", "user_api_key_end_user_id": "u1"}),
);
let expected = format!(
"{:x}",
Sha256::digest(b"model: gpt-4.1user_api_key: k1user_api_key_end_user_id: u1")
);
assert_eq!(
cache_key(&semantic_key(&end_user_request, "end_user")),
expected
);
let preset_request = native_request(
CacheKeyInput {
preset: Some("preset-key".to_owned()),
..Default::default()
},
json!({"user_api_key": "k1"}),
);
assert_eq!(
semantic_key(&preset_request, "end_user").preset.as_deref(),
Some("preset-key")
);
assert!(semantic_key(&preset_request, "end_user").fields.is_empty());
assert_eq!(preset_request.semantic().context.scope, None);
assert_eq!(
preset_request
.scoped_semantic("end_user")
.context
.scope
.as_deref(),
Some("end_user")
);
}
}

View file

@ -1,211 +0,0 @@
use crate::cache::cache_error;
use crate::execution::run_async;
use std::{collections::VecDeque, time::Duration};
use litellm_cache::Error;
use litellm_host_python::{Execution, ExecutionBody, ExecutionStep};
use pyo3::{
PyTraverseError, PyVisit,
exceptions::{PyException, PyRuntimeError},
prelude::*,
};
use serde_json::Value;
use super::{
backend::{NativeResponseCache, SemanticReply},
embedder::{PythonEmbedder, with_prepared_embedding},
request::{NativeRequest, now},
};
pub(super) enum SemanticOperation {
Lookup(NativeRequest),
/// A lookup that also reports the similarity, as `SemanticReply`.
LookupSemantic(NativeRequest),
Store(NativeRequest, Value),
StoreBatch(VecDeque<(NativeRequest, Value)>),
}
/// What an exception from the Python embedder means for the operation.
#[derive(Clone, Copy)]
pub(super) enum EmbeddingFailure {
/// Raise the Python exception unchanged.
Propagate,
/// Treat the embedding as unavailable and let the backend report that.
Unavailable,
}
enum Phase {
Start,
AwaitingEmbedding,
AwaitingBackend,
}
/// Runs a semantic cache operation whose embedding comes from Python: await the Python
/// embedder in the caller's event loop, seed the native backend with the vector, await the
/// backend, and repeat for each entry of a batch.
pub(super) struct SemanticExecution {
service: NativeResponseCache,
embedder: PythonEmbedder,
failure: EmbeddingFailure,
operation: SemanticOperation,
pending: Option<(NativeRequest, Option<Value>)>,
phase: Phase,
now: Duration,
}
impl SemanticExecution {
pub(super) fn new(
service: NativeResponseCache,
embedder: PythonEmbedder,
failure: EmbeddingFailure,
operation: SemanticOperation,
) -> Self {
Self {
service,
embedder,
failure,
operation,
pending: None,
phase: Phase::Start,
now: now(),
}
}
/// Takes the next entry of the operation; `None` once a batch is exhausted.
fn next_pending(&mut self) -> Option<(NativeRequest, Option<Value>)> {
match &mut self.operation {
SemanticOperation::Lookup(request) | SemanticOperation::LookupSemantic(request) => {
Some((request.clone(), None))
}
SemanticOperation::Store(request, response) => {
Some((request.clone(), Some(std::mem::take(response))))
}
SemanticOperation::StoreBatch(queue) => queue
.pop_front()
.map(|(request, response)| (request, Some(response))),
}
}
fn start(&mut self, py: Python<'_>) -> PyResult<ExecutionStep> {
let Some(pending) = self.next_pending() else {
return Ok(ExecutionStep::Return(py.None()));
};
let (request, response) = &pending;
let enabled = match response {
None => request.controls.reads(),
Some(_) => request.controls.writes(),
};
let input = enabled
.then(|| self.service.embedding_input(request))
.flatten();
self.pending = Some(pending);
let Some(input) = input else {
return self.backend_step(py, Err(Error::Unavailable));
};
let awaitable =
self.embedder
.async_embedding(py, &input.prompt, input.metadata.as_ref())?;
self.phase = Phase::AwaitingEmbedding;
Ok(ExecutionStep::Await(awaitable))
}
fn embedded(&mut self, py: Python<'_>, result: PyResult<Py<PyAny>>) -> PyResult<ExecutionStep> {
let seed = match result {
Ok(vector) => {
PythonEmbedder::extract(vector.into_bound(py)).map_err(|_| Error::Unavailable)
}
Err(error) => match self.embedding_failure() {
EmbeddingFailure::Propagate => return Err(error),
EmbeddingFailure::Unavailable if error.is_instance_of::<PyException>(py) => {
Err(Error::Unavailable)
}
EmbeddingFailure::Unavailable => return Err(error),
},
};
self.backend_step(py, seed)
}
/// Python's semantic lookups catch embedding errors and stamp a similarity of `0.0`.
fn embedding_failure(&self) -> EmbeddingFailure {
match self.operation {
SemanticOperation::LookupSemantic(_) => EmbeddingFailure::Unavailable,
_ => self.failure,
}
}
fn backend_step(
&mut self,
py: Python<'_>,
seed: Result<Vec<f32>, Error>,
) -> PyResult<ExecutionStep> {
self.phase = Phase::AwaitingBackend;
let (request, response) = self.pending.take().ok_or_else(|| {
PyRuntimeError::new_err("semantic execution resumed without a pending operation")
})?;
let service = self.service.clone();
let now = self.now;
let with_similarity = matches!(self.operation, SemanticOperation::LookupSemantic(_));
let future = async move {
match response {
None if with_similarity => service
.async_lookup_semantic(&request, now)
.await
.map(|lookup| Reply::Semantic(lookup.into())),
None => service.async_lookup(&request, now).await.map(Reply::Plain),
Some(response) => service
.async_store(&request, response, now)
.await
.map(|_| Reply::Plain(None)),
}
};
let awaitable = run_async(py, with_prepared_embedding(seed, future), cache_error)?;
Ok(ExecutionStep::Await(awaitable.unbind()))
}
fn resume_py(
&mut self,
py: Python<'_>,
result: Option<PyResult<Py<PyAny>>>,
) -> PyResult<ExecutionStep> {
match (&self.phase, result) {
(Phase::Start, None) => self.start(py),
(Phase::AwaitingEmbedding, Some(result)) => self.embedded(py, result),
(Phase::AwaitingBackend, Some(Err(error))) => Err(error),
(Phase::AwaitingBackend, Some(Ok(value))) => {
let more = matches!(
&self.operation,
SemanticOperation::StoreBatch(queue) if !queue.is_empty()
);
if more {
self.phase = Phase::Start;
return self.start(py);
}
Ok(ExecutionStep::Return(value))
}
_ => Err(PyRuntimeError::new_err(
"invalid semantic cache execution state",
)),
}
}
}
#[derive(serde::Serialize)]
#[serde(untagged)]
enum Reply {
Plain(Option<Value>),
Semantic(SemanticReply),
}
impl ExecutionBody for SemanticExecution {
fn resume(&mut self, result: Option<PyResult<Py<PyAny>>>) -> PyResult<ExecutionStep> {
Python::attach(|py| self.resume_py(py, result))
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
self.embedder.traverse(visit)
}
}
pub(super) fn drive(py: Python<'_>, body: SemanticExecution) -> PyResult<Bound<'_, PyAny>> {
Execution::new(body, crate::lifecycle::binding).into_coroutine(py)
}

View file

@ -1,5 +1,8 @@
use crate::cache::cache_error;
use std::{sync::Arc, time::Duration};
use std::{
sync::Arc,
time::{Duration, SystemTime, UNIX_EPOCH},
};
use litellm_cache::{DeleteCache, DisconnectCache, PingCache};
use litellm_host_python::{from_py, release_gil, to_py};
@ -118,8 +121,8 @@ impl NativeCacheHandle {
fn get(&self, py: Python<'_>, key: String) -> PyResult<Py<PyAny>> {
self.check_process()?;
let request = request(key, None)?;
let value = release_gil(py, || self.backend.lookup(&request, super::request::now()))
.map_err(cache_error)?;
let value =
release_gil(py, || self.backend.lookup(&request, now())).map_err(cache_error)?;
to_py(py, &value)
}
@ -134,10 +137,7 @@ impl NativeCacheHandle {
self.check_process()?;
let request = request(key, ttl)?;
let value: Value = from_py(value)?;
release_gil(py, || {
self.backend.store(&request, value, super::request::now())
})
.map_err(cache_error)
release_gil(py, || self.backend.store(&request, value, now())).map_err(cache_error)
}
fn async_get<'py>(&self, py: Python<'py>, key: String) -> PyResult<Bound<'py, PyAny>> {
@ -146,7 +146,7 @@ impl NativeCacheHandle {
let backend = self.backend.clone();
crate::execution::run_async(
py,
async move { backend.async_lookup(&request, super::request::now()).await },
async move { backend.async_lookup(&request, now()).await },
cache_error,
)
}
@ -165,11 +165,7 @@ impl NativeCacheHandle {
let backend = self.backend.clone();
crate::execution::run_async(
py,
async move {
backend
.async_store(&request, value, super::request::now())
.await
},
async move { backend.async_store(&request, value, now()).await },
cache_error,
)
}
@ -190,11 +186,7 @@ impl NativeCacheHandle {
let backend = self.backend.clone();
crate::execution::run_async(
py,
async move {
backend
.async_store_batch(entries, super::request::now())
.await
},
async move { backend.async_store_batch(entries, now()).await },
cache_error,
)
}
@ -345,3 +337,9 @@ pub(in crate::cache) fn configured(
},
))
}
fn now() -> Duration {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
}

View file

@ -1,6 +1,6 @@
# Python cache delegation
This directory lets Rust inference use a selected Python cache. `service.rs` implements the injected Rust response-cache service and yields typed cache operations. `host.rs` calls the Python cache's sync or async API and delivers the result back to Rust. `callback.rs` provides Python cache delegation for the Python-facing cache runtime
This directory lets Rust inference use a selected Python cache. `service.rs` implements the injected Rust response-cache service and yields typed cache operations. `host.rs` calls the Python cache's sync or async API and delivers the result back to Rust.
Keep the shared inference protocol wrapper and adapter selection in the parent module. Receive the configured cache and prepared arguments from the parent module. Do not discover global configuration, choose native backends, or move inference to Python. Core remains independent of Python objects and cache implementation details

View file

@ -1,165 +0,0 @@
use pyo3::{
PyTraverseError, PyVisit,
exceptions::{PyTypeError, PyValueError},
prelude::*,
types::{PyDict, PyList, PyTuple},
};
use crate::cache::future::ready_none;
pub(in crate::cache) struct PythonCallback(Py<PyAny>);
impl PythonCallback {
pub(in crate::cache) fn new(object: Py<PyAny>) -> Self {
Self(object)
}
pub(in crate::cache) fn lookup<'py>(
&self,
py: Python<'py>,
kwargs: Option<&Bound<'py, PyDict>>,
) -> PyResult<Bound<'py, PyAny>> {
self.0
.bind(py)
.call_method("get_cache", (), Some(callback_kwargs(kwargs)?))
}
pub(in crate::cache) fn async_lookup<'py>(
&self,
py: Python<'py>,
kwargs: Option<&Bound<'py, PyDict>>,
) -> PyResult<Bound<'py, PyAny>> {
self.0
.bind(py)
.call_method("async_get_cache", (), Some(callback_kwargs(kwargs)?))
}
pub(in crate::cache) fn store(
&self,
py: Python<'_>,
response: &Bound<'_, PyAny>,
kwargs: Option<&Bound<'_, PyDict>>,
) -> PyResult<()> {
self.0
.bind(py)
.call_method("add_cache", (response,), Some(callback_kwargs(kwargs)?))
.map(|_| ())
}
pub(in crate::cache) fn async_store<'py>(
&self,
py: Python<'py>,
response: &Bound<'py, PyAny>,
kwargs: Option<&Bound<'py, PyDict>>,
) -> PyResult<Bound<'py, PyAny>> {
self.0.bind(py).call_method(
"async_add_cache",
(response,),
Some(callback_kwargs(kwargs)?),
)
}
pub(in crate::cache) fn lookup_batch<'py>(
&self,
py: Python<'py>,
requests: &Bound<'py, PyAny>,
kwargs: Option<&Bound<'py, PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let results = PyList::empty(py);
for kwargs in batch_callback_kwargs(requests, kwargs)? {
results.append(
self.0
.bind(py)
.call_method("get_cache", (), Some(&kwargs))?,
)?;
}
Ok(results.into_any())
}
pub(in crate::cache) fn async_lookup_batch<'py>(
&self,
py: Python<'py>,
requests: &Bound<'py, PyAny>,
kwargs: Option<&Bound<'py, PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let awaitables = batch_callback_kwargs(requests, kwargs)?
.iter()
.map(|kwargs| {
self.0
.bind(py)
.call_method("async_get_cache", (), Some(kwargs))
})
.collect::<PyResult<Vec<_>>>()?;
py.import("asyncio")?
.call_method1("gather", PyTuple::new(py, awaitables)?)
}
pub(in crate::cache) fn async_store_batch<'py>(
&self,
py: Python<'py>,
result: Option<&Bound<'py, PyAny>>,
kwargs: Option<&Bound<'py, PyDict>>,
) -> PyResult<Bound<'py, PyAny>> {
let result = result.ok_or_else(|| {
PyTypeError::new_err("Python cache callbacks require their original callback_result")
})?;
self.0.bind(py).call_method(
"async_add_cache_pipeline",
(result,),
Some(callback_kwargs(kwargs)?),
)
}
pub(in crate::cache) fn async_flush<'py>(
&self,
py: Python<'py>,
) -> PyResult<Bound<'py, PyAny>> {
let object = self.0.bind(py);
let backend = match object.getattr_opt("cache")? {
Some(backend) if !backend.is_none() => backend,
_ => object.clone(),
};
if backend.hasattr("async_flush_cache")? {
return backend.call_method0("async_flush_cache");
}
backend.call_method0("flush_cache")?;
ready_none(py)
}
pub(in crate::cache) fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
self.0.bind(py).call_method0("ping")
}
pub(in crate::cache) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.0)
}
}
fn callback_kwargs<'a, 'py>(
kwargs: Option<&'a Bound<'py, PyDict>>,
) -> PyResult<&'a Bound<'py, PyDict>> {
kwargs.ok_or_else(|| {
PyTypeError::new_err("Python cache callbacks require their original callback_kwargs")
})
}
fn batch_callback_kwargs<'py>(
requests: &Bound<'py, PyAny>,
kwargs: Option<&Bound<'py, PyAny>>,
) -> PyResult<Vec<Bound<'py, PyDict>>> {
let kwargs = kwargs
.ok_or_else(|| {
PyTypeError::new_err(
"Python cache callbacks require one original callback_kwargs mapping per request",
)
})?
.try_iter()?
.map(|item| Ok(item?.cast_into::<PyDict>()?))
.collect::<PyResult<Vec<_>>>()?;
if kwargs.len() != requests.len()? {
return Err(PyValueError::new_err(
"batch cache requests and callback_kwargs must have equal lengths",
));
}
Ok(kwargs)
}

View file

@ -1,8 +1,6 @@
mod callback;
mod host;
mod service;
pub(super) use callback::PythonCallback;
pub(crate) use host::PythonCache;
pub(crate) use service::CacheCall;
pub(super) use service::service;

View file

@ -1,385 +0,0 @@
use crate::execution::run_async;
use litellm_cache_response::PartialHits;
use litellm_host_python::{ExecutionStep, from_py, release_gil, to_py};
use pyo3::{
PyTraverseError, PyVisit,
exceptions::{PyRuntimeError, PyValueError},
prelude::*,
types::PyDict,
};
use serde_json::Value;
use super::{
cache_error,
future::{ready_none, ready_value},
native::activation::activate,
native::backend::{NativeResponseCache, SemanticReply},
native::config::{CacheConfigProjection, NativeCacheConfig},
native::request::{now, request, requests},
python::PythonCallback,
};
use crate::errors::RustBridgeDeclined;
pub(super) enum CacheBinding {
Disabled,
Native(NativeResponseCache),
PythonCallback(PythonCallback),
}
#[pyclass(frozen, name = "_ResponseCacheRuntime")]
pub(crate) struct ResolvedCache {
binding: CacheBinding,
guard: Option<super::native::facade::FacadeGuard>,
pid: u32,
}
impl ResolvedCache {
pub(super) fn new(binding: CacheBinding) -> Self {
Self {
binding,
guard: None,
pid: std::process::id(),
}
}
pub(super) fn with_guard(mut self, guard: super::native::facade::FacadeGuard) -> Self {
self.guard = Some(guard);
self
}
pub(super) fn native_service(&self) -> PyResult<Option<NativeResponseCache>> {
self.check_process()?;
Ok(match &self.binding {
CacheBinding::Native(service) => Some(service.clone()),
_ => None,
})
}
fn check_process(&self) -> PyResult<()> {
if matches!(self.binding, CacheBinding::Native(_)) && self.pid != std::process::id() {
return Err(PyRuntimeError::new_err(
"native cache bindings must be resolved again after fork",
));
}
Ok(())
}
pub(crate) fn lookup_step(
&self,
py: Python<'_>,
input: &Bound<'_, PyAny>,
kwargs: Option<&Bound<'_, PyDict>>,
) -> PyResult<ExecutionStep> {
self.check_process()?;
let awaitable = match &self.binding {
CacheBinding::Disabled => ready_none(py)?,
CacheBinding::Native(service) => {
let request = request(input)?;
service.async_lookup_py(py, request)?
}
CacheBinding::PythonCallback(callback) => callback.async_lookup(py, kwargs)?,
};
Ok(ExecutionStep::Await(awaitable.unbind()))
}
}
#[pymethods]
impl ResolvedCache {
#[staticmethod]
pub(crate) fn from_selected(cache: &Bound<'_, PyAny>) -> PyResult<Self> {
let py = cache.py();
let binding = if cache.is_none() {
CacheBinding::Disabled
} else if let Some(runtime) = cache
.getattr_opt("_native_cache")?
.filter(|value| !value.is_none())
{
let resolved = runtime
.getattr("native")?
.extract::<PyRef<'_, ResolvedCache>>()?;
match resolved.native_service()? {
Some(service) => {
if !resolved
.guard
.as_ref()
.is_some_and(|guard| guard.matches(py, cache).unwrap_or(false))
{
return Err(RustBridgeDeclined::new_err(
"native cache runtime no longer matches its facade",
));
}
CacheBinding::Native(service)
}
None => CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind())),
}
} else {
CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind()))
};
Ok(Self::new(binding))
}
#[staticmethod]
fn from_cache(cache: &Bound<'_, PyAny>) -> PyResult<Self> {
let config = match NativeCacheConfig::project(cache)? {
CacheConfigProjection::Native(config) => *config,
CacheConfigProjection::Unsupported(reason) => {
return Err(RustBridgeDeclined::new_err(reason.message()));
}
};
let backend = cache.getattr("cache")?;
let service = activate(cache.py(), &backend, config)?;
let resolved = Self::new(CacheBinding::Native(service.clone()));
Ok(
match super::native::facade::FacadeGuard::capture(cache.py(), cache, &service) {
Ok(guard) => resolved.with_guard(guard),
Err(_) => resolved,
},
)
}
#[getter]
fn kind(&self) -> &'static str {
match self.binding {
CacheBinding::Disabled => "disabled",
CacheBinding::Native(_) => "native",
CacheBinding::PythonCallback(_) => "python_callback",
}
}
#[pyo3(signature = (request, *, callback_kwargs=None))]
fn lookup(
&self,
py: Python<'_>,
request: &Bound<'_, PyAny>,
callback_kwargs: Option<&Bound<'_, PyDict>>,
) -> PyResult<Py<PyAny>> {
self.check_process()?;
match &self.binding {
CacheBinding::Disabled => Ok(py.None()),
CacheBinding::Native(service) => {
let request = self::request(request)?;
let service = service.clone();
let response = release_gil(py, move || service.lookup(&request, now()))
.map_err(cache_error)?;
to_py(py, &response)
}
CacheBinding::PythonCallback(callback) => {
callback.lookup(py, callback_kwargs).map(Bound::unbind)
}
}
}
/// `(response, similarity)`: the similarity is `None` when the backend reports none.
fn lookup_semantic(&self, py: Python<'_>, request: &Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
self.check_process()?;
match &self.binding {
CacheBinding::Native(service) => {
let request = self::request(request)?;
let service = service.clone();
let lookup = release_gil(py, move || service.lookup_semantic(&request, now()))
.map_err(cache_error)?;
to_py(py, &SemanticReply::from(lookup))
}
CacheBinding::Disabled => to_py(py, &SemanticReply(None, None)),
CacheBinding::PythonCallback(_) => Err(PyRuntimeError::new_err(
"semantic lookups require a native cache binding",
)),
}
}
fn async_lookup_semantic<'py>(
&self,
py: Python<'py>,
request: &Bound<'py, PyAny>,
) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
match &self.binding {
CacheBinding::Native(service) => {
service.async_lookup_semantic_py(py, self::request(request)?)
}
CacheBinding::Disabled => ready_value(py, &SemanticReply(None, None)),
CacheBinding::PythonCallback(_) => Err(PyRuntimeError::new_err(
"semantic lookups require a native cache binding",
)),
}
}
#[pyo3(signature = (request, response, *, callback_kwargs=None))]
fn store(
&self,
py: Python<'_>,
request: &Bound<'_, PyAny>,
response: &Bound<'_, PyAny>,
callback_kwargs: Option<&Bound<'_, PyDict>>,
) -> PyResult<()> {
self.check_process()?;
match &self.binding {
CacheBinding::Disabled => Ok(()),
CacheBinding::Native(service) => {
let request = self::request(request)?;
let response: Value = from_py(response)?;
let service = service.clone();
release_gil(py, move || service.store(&request, response, now()))
.map_err(cache_error)
}
CacheBinding::PythonCallback(callback) => callback.store(py, response, callback_kwargs),
}
}
#[pyo3(signature = (requests, *, callback_kwargs=None))]
fn lookup_batch(
&self,
py: Python<'_>,
requests: &Bound<'_, PyAny>,
callback_kwargs: Option<&Bound<'_, PyAny>>,
) -> PyResult<Py<PyAny>> {
self.check_process()?;
match &self.binding {
CacheBinding::Disabled => {
let requests = self::requests(requests)?;
to_py(py, &PartialHits::new(vec![None; requests.len()]))
}
CacheBinding::Native(service) => {
let requests = self::requests(requests)?;
let service = service.clone();
let response = release_gil(py, move || service.lookup_batch(&requests, now()))
.map_err(cache_error)?;
to_py(py, &response)
}
CacheBinding::PythonCallback(callback) => callback
.lookup_batch(py, requests, callback_kwargs)
.map(Bound::unbind),
}
}
#[pyo3(signature = (request, *, callback_kwargs=None))]
fn async_lookup<'py>(
&self,
py: Python<'py>,
request: &Bound<'py, PyAny>,
callback_kwargs: Option<&Bound<'py, PyDict>>,
) -> PyResult<Bound<'py, PyAny>> {
let ExecutionStep::Await(awaitable) = self.lookup_step(py, request, callback_kwargs)?
else {
unreachable!()
};
Ok(awaitable.into_bound(py))
}
#[pyo3(signature = (request, response, *, callback_kwargs=None))]
fn async_store<'py>(
&self,
py: Python<'py>,
request: &Bound<'py, PyAny>,
response: &Bound<'py, PyAny>,
callback_kwargs: Option<&Bound<'py, PyDict>>,
) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
match &self.binding {
CacheBinding::Disabled => ready_none(py),
CacheBinding::Native(service) => {
let request = self::request(request)?;
let response: Value = from_py(response)?;
service.async_store_py(py, request, response)
}
CacheBinding::PythonCallback(callback) => {
callback.async_store(py, response, callback_kwargs)
}
}
}
#[pyo3(signature = (requests, *, callback_kwargs=None))]
fn async_lookup_batch<'py>(
&self,
py: Python<'py>,
requests: &Bound<'py, PyAny>,
callback_kwargs: Option<&Bound<'py, PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
match &self.binding {
CacheBinding::Disabled => {
let requests = self::requests(requests)?;
ready_value(py, &PartialHits::new(vec![None; requests.len()]))
}
CacheBinding::Native(service) => {
let requests = self::requests(requests)?;
let service = service.clone();
run_async(
py,
async move { service.async_lookup_batch(&requests, now()).await },
cache_error,
)
}
CacheBinding::PythonCallback(callback) => {
callback.async_lookup_batch(py, requests, callback_kwargs)
}
}
}
#[pyo3(signature = (requests, responses, *, callback_result=None, callback_kwargs=None))]
fn async_store_batch<'py>(
&self,
py: Python<'py>,
requests: &Bound<'py, PyAny>,
responses: &Bound<'py, PyAny>,
callback_result: Option<&Bound<'py, PyAny>>,
callback_kwargs: Option<&Bound<'py, PyDict>>,
) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
match &self.binding {
CacheBinding::Disabled => ready_none(py),
CacheBinding::Native(service) => {
let requests = self::requests(requests)?;
let responses: Vec<Value> = from_py(responses)?;
if requests.len() != responses.len() {
return Err(PyValueError::new_err(
"batch cache requests and responses must have equal lengths",
));
}
let entries = requests.into_iter().zip(responses).collect();
service.async_store_batch_py(py, entries)
}
CacheBinding::PythonCallback(callback) => {
callback.async_store_batch(py, callback_result, callback_kwargs)
}
}
}
fn async_flush<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
match &self.binding {
CacheBinding::Disabled => ready_none(py),
CacheBinding::Native(service) => {
let service = service.clone();
run_async(py, async move { service.async_flush().await }, cache_error)
}
CacheBinding::PythonCallback(callback) => callback.async_flush(py),
}
}
fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
self.check_process()?;
match &self.binding {
CacheBinding::Disabled => ready_none(py),
CacheBinding::Native(service) => {
let service = service.clone();
run_async(
py,
async move { service.test_connection().await },
cache_error,
)
}
CacheBinding::PythonCallback(callback) => callback.ping(py),
}
}
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
if let CacheBinding::PythonCallback(callback) = &self.binding {
callback.traverse(&visit)?;
}
if let Some(guard) = &self.guard {
guard.traverse(visit)?;
}
Ok(())
}
}

View file

@ -17,7 +17,6 @@ mod tokenizer;
#[pymodule(gil_used = true)]
mod _native {
use crate::cache::ResolvedCache;
#[cfg(feature = "panic-test")]
#[pymodule_export]
use crate::diagnostics::_panic_for_test;
@ -64,7 +63,6 @@ mod _native {
"NativeCacheHandle",
py.get_type::<crate::cache::NativeCacheHandle>(),
)?;
dict.set_item("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())?;
dict.set_item(
"_SecretManagerRuntime",
py.get_type::<crate::secrets::runtime::NativeSecretManager>(),

View file

@ -16,8 +16,7 @@ import traceback
from collections.abc import Generator, Mapping
from contextlib import contextmanager
from enum import Enum
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
from typing import Any, Final, Literal
from pydantic import BaseModel
@ -42,26 +41,6 @@ from .redis_cluster_cache import RedisClusterCache
from .redis_semantic_cache import RedisSemanticCache
from .s3_cache import S3Cache
if TYPE_CHECKING:
from litellm.rust_bridge.response_cache import NativeCacheRequest, ResponseCacheRuntime
def _native_response(result: object) -> object:
"""The value Python's own reader would return for `result` once it is cached.
Python stores a model as its JSON text and `json.loads` a string response on read, so the
native store receives the decoded value and writes the envelope shape Python reads.
"""
if isinstance(result, BaseModel):
return json.loads(result.model_dump_json())
if isinstance(result, str):
try:
return json.loads(result)
except ValueError:
return result
return result
RESPONSE_CACHE_TARGET: Final = "llm_response"
@ -93,8 +72,6 @@ class CacheMode(str, Enum):
#### LiteLLM.Completion / Embedding Cache ####
class Cache:
_native_cache: "ResponseCacheRuntime | None" = None
def __init__(
self,
type: LiteLLMCacheType | None = LiteLLMCacheType.LOCAL,
@ -620,13 +597,6 @@ class Cache:
if "semantic-similarity" in cache_lookup_metadata:
original_metadata["semantic-similarity"] = cache_lookup_metadata["semantic-similarity"]
@staticmethod
def _stamp_semantic_similarity(kwargs: Mapping[str, object], similarity: float | None) -> None:
"""Write a native semantic lookup's similarity where the Python backends put it."""
metadata: Final = kwargs.get("metadata")
if similarity is not None and isinstance(metadata, dict):
metadata["semantic-similarity"] = similarity
def get_cache(self, dynamic_cache_object: BaseCache | None = None, **kwargs):
"""
Retrieves the cached result for the given arguments.
@ -646,15 +616,6 @@ class Cache:
cache_key = kwargs["cache_key"]
else:
cache_key = self.get_cache_key(**kwargs)
if cache_key is not None and self._native_cache is not None:
request = self._native_cache.request(self, MappingProxyType({**kwargs, "cache_key": cache_key}))
if request is None:
return None
if not self._is_semantic_cache():
return self._native_cache.lookup(request)
response, similarity = self._native_cache.lookup_semantic(request)
self._stamp_semantic_similarity(kwargs, similarity)
return response
if cache_key is not None:
cache_control_args: Final[DynamicCacheControl] = kwargs.get("cache", {})
max_age = cache_control_args.get("s-maxage") or cache_control_args.get("s-max-age") or float("inf")
@ -688,15 +649,6 @@ class Cache:
cache_key = kwargs["cache_key"]
else:
cache_key = self.get_cache_key(**kwargs)
if cache_key is not None and self._native_cache is not None:
request = self._native_cache.request(self, MappingProxyType({**kwargs, "cache_key": cache_key}))
if request is None:
return None
if not self._is_semantic_cache():
return await self._native_cache.async_lookup(request)
response, similarity = await self._native_cache.async_lookup_semantic(request)
self._stamp_semantic_similarity(kwargs, similarity)
return response
if cache_key is not None:
cache_control_args: Final = kwargs.get("cache", {})
max_age: Final = cache_control_args.get(
@ -756,11 +708,6 @@ class Cache:
if self.should_use_cache(**kwargs) is not True:
return
with response_cache_phase("set"):
if self._native_cache is not None:
request = self._native_request(kwargs)
if request is not None:
self._native_cache.store(request, _native_response(result))
return
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
self.cache.set_cache(cache_key, cached_data, **kwargs)
except Exception as e:
@ -781,11 +728,6 @@ class Cache:
if self.should_use_cache(**kwargs) is not True:
return
with response_cache_phase("set"):
if self._native_cache is not None:
request = self._native_request(kwargs)
if request is not None:
await self._native_cache.async_store(request, _native_response(result))
return
if self.type == "redis" and self.redis_flush_size is not None:
# high traffic - fill in results in memory and then flush
await self.batch_cache_write(result, **kwargs)
@ -970,35 +912,13 @@ class Cache:
)
cache_list.append((cache_key, cached_data))
if self._native_cache is not None:
entries: Final = tuple(
(request, cached_data["response"])
for cache_key, cached_data in cache_list
if (request := self._native_request(MappingProxyType({**kwargs, "cache_key": cache_key})))
is not None
)
await self._native_cache.async_store_batch(
tuple(request for request, _ in entries),
tuple(response for _, response in entries),
)
elif dynamic_cache_object is not None:
if dynamic_cache_object is not None:
await dynamic_cache_object.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
else:
await self.cache.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
except Exception as e:
self._log_add_cache_failure(e)
def _native_request(self, kwargs: Mapping[str, object]) -> "NativeCacheRequest | None":
if self._native_cache is None:
return None
cache_key: Final = kwargs.get("cache_key")
return self._native_cache.request(
self,
kwargs
if isinstance(cache_key, str)
else MappingProxyType({**kwargs, "cache_key": self.get_cache_key(**kwargs)}),
)
def should_use_cache(self, **kwargs):
"""
Returns true if we should use the cache for LLM API calls

View file

@ -179,65 +179,6 @@ class ResponsesWebSocketConnection:
def recv_text(self) -> Future[str | None]: ...
def close(self) -> Future[None]: ...
@final
class _ResponseCacheRuntime:
@staticmethod
def from_cache(cache: object) -> _ResponseCacheRuntime: ...
@staticmethod
def from_selected(cache: object) -> _ResponseCacheRuntime: ...
@property
def kind(self) -> str: ...
def lookup(
self,
request: object,
*,
callback_kwargs: Mapping[str, object] | Sequence[object] | None = None,
) -> object: ...
def lookup_semantic(self, request: object) -> tuple[object, float | None]: ...
def store(
self,
request: object,
response: object,
*,
callback_kwargs: Mapping[str, object] | None = None,
) -> None: ...
def lookup_batch(
self,
requests: Sequence[object],
*,
callback_kwargs: Sequence[object] | None = None,
) -> object: ...
def async_lookup(
self,
request: object,
*,
callback_kwargs: Mapping[str, object] | None = None,
) -> Future[object]: ...
def async_lookup_semantic(self, request: object) -> Future[tuple[object, float | None]]: ...
def async_store(
self,
request: object,
response: object,
*,
callback_kwargs: Mapping[str, object] | None = None,
) -> Future[None]: ...
def async_lookup_batch(
self,
requests: Sequence[object],
*,
callback_kwargs: Sequence[object] | None = None,
) -> Future[object]: ...
def async_store_batch(
self,
requests: Sequence[object],
responses: Sequence[object],
*,
callback_result: object = None,
callback_kwargs: Mapping[str, object] | None = None,
) -> Future[object]: ...
def async_flush(self) -> Future[None]: ...
def ping(self) -> Future[object]: ...
@final
class TokenCounter:
@staticmethod

View file

@ -1,146 +0,0 @@
from __future__ import annotations
import math
from collections.abc import Awaitable, Mapping, Sequence
from dataclasses import dataclass
from typing import Final, Protocol, cast
from typing_extensions import ReadOnly, Required, TypedDict
class CacheFacade(Protocol):
@property
def type(self) -> object: ...
@property
def ttl(self) -> float | None: ...
@property
def semantic_cache_scope(self) -> str: ...
def get_cache_key(self, **kwargs: object) -> str: ... # kwargs-ok: mirrors the legacy cache facade contract
class NativeCacheKey(TypedDict):
preset: ReadOnly[str]
class NativeCacheRequest(TypedDict, total=False):
key: Required[ReadOnly[NativeCacheKey]]
ttl_seconds: ReadOnly[float | None]
max_age_seconds: ReadOnly[float | None]
messages: ReadOnly[object | None]
input: ReadOnly[object | None]
metadata: ReadOnly[object | None]
litellm_metadata: ReadOnly[object | None]
litellm_params: ReadOnly[object | None]
scope: ReadOnly[str]
class NativeResponseCacheRuntime(Protocol):
@property
def kind(self) -> str: ...
def lookup(self, request: NativeCacheRequest) -> object: ...
def lookup_semantic(self, request: NativeCacheRequest) -> tuple[object, float | None]: ...
def store(self, request: NativeCacheRequest, response: object) -> None: ...
def lookup_batch(self, requests: Sequence[NativeCacheRequest]) -> object: ...
def async_lookup(self, request: NativeCacheRequest) -> Awaitable[object]: ...
def async_lookup_semantic(self, request: NativeCacheRequest) -> Awaitable[tuple[object, float | None]]: ...
def async_store(self, request: NativeCacheRequest, response: object) -> Awaitable[None]: ...
def async_lookup_batch(self, requests: Sequence[NativeCacheRequest]) -> Awaitable[object]: ...
def async_store_batch(
self,
requests: Sequence[NativeCacheRequest],
responses: Sequence[object],
) -> Awaitable[object]: ...
def async_flush(self) -> Awaitable[None]: ...
def ping(self) -> Awaitable[object]: ...
@dataclass(frozen=True, slots=True)
class ResponseCacheRuntime:
native: NativeResponseCacheRuntime
@property
def kind(self) -> str:
return self.native.kind
def request(self, cache: CacheFacade, kwargs: Mapping[str, object]) -> NativeCacheRequest | None:
key_value: Final = kwargs.get("cache_key")
key: Final = key_value if isinstance(key_value, str) else cache.get_cache_key(**dict(kwargs))
if not key:
return None
control_value: Final = kwargs.get("cache")
control: Final = _string_mapping(control_value)
configured_ttl: Final = cache.ttl if cache.ttl is not None else _duration(kwargs.get("ttl"))
control_ttl: Final = _duration(control.get("ttl"))
current_max_age: Final = _duration(control.get("s-max-age"))
legacy_max_age: Final = _duration(control.get("s-maxage"))
ttl: Final = configured_ttl if control_ttl is None else control_ttl
max_age: Final = legacy_max_age if current_max_age is None else current_max_age
return NativeCacheRequest(
key=NativeCacheKey(preset=key),
ttl_seconds=ttl,
max_age_seconds=max_age,
messages=kwargs.get("messages"),
input=kwargs.get("input"),
metadata=kwargs.get("metadata"),
litellm_metadata=kwargs.get("litellm_metadata"),
litellm_params=kwargs.get("litellm_params"),
scope=cache.semantic_cache_scope,
)
def lookup(self, request: NativeCacheRequest) -> object:
return self.native.lookup(request)
def lookup_semantic(self, request: NativeCacheRequest) -> tuple[object, float | None]:
"""The cached response and the similarity a semantic backend reports, if any."""
response, similarity = self.native.lookup_semantic(request)
return response, similarity
def store(self, request: NativeCacheRequest, response: object) -> None:
self.native.store(request, response)
def lookup_batch(self, requests: Sequence[NativeCacheRequest]) -> object:
return self.native.lookup_batch(requests)
async def async_lookup(self, request: NativeCacheRequest) -> object:
return await self.native.async_lookup(request)
async def async_lookup_semantic(self, request: NativeCacheRequest) -> tuple[object, float | None]:
response, similarity = await self.native.async_lookup_semantic(request)
return response, similarity
async def async_store(self, request: NativeCacheRequest, response: object) -> None:
await self.native.async_store(request, response)
async def async_lookup_batch(self, requests: Sequence[NativeCacheRequest]) -> object:
return await self.native.async_lookup_batch(requests)
async def async_store_batch(
self,
requests: Sequence[NativeCacheRequest],
responses: Sequence[object],
) -> object:
return await self.native.async_store_batch(requests, responses)
async def ping(self) -> object:
return await self.native.ping()
async def async_flush(self) -> None:
await self.native.async_flush()
def _duration(value: object) -> float | None:
if isinstance(value, bool) or not isinstance(value, int | float):
return None
duration: Final = float(value)
return duration if math.isfinite(duration) and duration >= 0 else None
def _string_mapping(value: object) -> Mapping[str, object]:
if not isinstance(value, Mapping):
return {}
source: Final = cast(Mapping[object, object], value)
return {key: item for key, item in source.items() if isinstance(key, str)}

View file

@ -5,8 +5,6 @@ from typing import Final
import fakeredis
import pytest
from tests.test_litellm_rust.support.s3_stub import S3Stub
@pytest.fixture
def redis_url() -> Generator[str]:
@ -19,12 +17,3 @@ def redis_url() -> Generator[str]:
server.shutdown()
server.server_close()
worker.join(timeout=5)
@pytest.fixture
def s3_stub() -> Generator[S3Stub]:
stub: Final = S3Stub()
try:
yield stub
finally:
stub.close()

View file

@ -1,165 +0,0 @@
import asyncio
import json
import os
import time
import uuid
from collections.abc import Generator
from types import SimpleNamespace
from typing import Final, cast
import pytest
from azure.storage.blob import ContainerClient
from litellm.caching.azure_blob_cache import AzureBlobCache
from litellm.caching.caching import Cache
from litellm.rust_bridge import _native
from litellm.types.caching import LiteLLMCacheType
from tests.test_litellm_rust.support.cache import (
CacheLookup,
CacheTestResolver,
activate_native,
assert_native_runtime,
completion_kwargs,
request,
)
from tests.test_litellm_rust.support.isolation import rebound
pytestmark: Final = pytest.mark.requires_rust_extension
@pytest.fixture
def azure_blob_facade() -> Generator[Cache]:
account_url: Final = os.environ.get("AZURE_BLOB_CACHE_ACCOUNT_URL")
if account_url is None:
pytest.skip(
"live Azure Blob parity needs AZURE_BLOB_CACHE_ACCOUNT_URL plus DefaultAzureCredential inputs in the environment"
)
facade: Final = Cache(
type=LiteLLMCacheType.AZURE_BLOB,
azure_account_url=account_url,
azure_blob_container=f"litellm-parity-{uuid.uuid4().hex[:12]}",
)
backend: Final = facade.cache
assert isinstance(backend, AzureBlobCache)
try:
yield facade
finally:
backend.container_client.delete_container()
asyncio.run(backend.disconnect())
def test_azure_blob_facade_serves_natively_and_python_reads_the_same_blobs(azure_blob_facade: Cache) -> None:
backend: Final = azure_blob_facade.cache
assert isinstance(backend, AzureBlobCache)
activate_native(azure_blob_facade)
account_url: Final = backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}")
resolver: Final = CacheTestResolver(SimpleNamespace(cache=azure_blob_facade))
native: Final = resolver.resolve()
assert native.kind == "native"
response: Final = {
"choices": [{"text": "caf\u00e9 \u2603"}],
"usage": {"total_tokens": 3},
"flag": True,
"empty": None,
}
native.store({**request("sync"), "ttl_seconds": 0.001}, response)
native.store(request("sync"), {"choices": [{"text": "second"}]})
time.sleep(0.01)
stored: Final = json.loads(backend.container_client.download_blob("sync").readall())
assert stored["response"] == response
assert isinstance(stored["timestamp"], float)
assert native.lookup(request("sync")) == response
with rebound(azure_blob_facade, "_native_cache", None):
assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response
backend.set_cache("python", {"timestamp": time.time(), "response": response})
backend.set_cache("legacy", "bare legacy value")
backend.container_client.upload_blob("invalid", b"{not json", overwrite=True)
assert native.lookup(request("python")) == response
with rebound(azure_blob_facade, "_native_cache", None):
assert native.lookup(request("legacy")) == cast(CacheLookup, azure_blob_facade).get_cache(cache_key="legacy")
assert native.lookup_batch([request("python"), request("missing"), request("invalid"), request("sync")]) == {
"values": [response, None, None, response],
"missing_indices": [1, 2],
}
with rebound(azure_blob_facade, "ttl", 12):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
with rebound(backend, "container_client", ContainerClient.from_container_url(backend.container_client.url)):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
def custom_get(*_args: object, **_kwargs: object) -> None:
return None
with rebound(backend, "get_cache", custom_get):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
with rebound(azure_blob_facade, "_native_cache", None):
assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response
class CustomBlobCache(AzureBlobCache):
pass
with rebound(azure_blob_facade, "cache", CustomBlobCache(account_url, backend.container_client.container_name)):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
async def test_azure_blob_native_async_writes_overwrite_batch_and_flush_like_python(azure_blob_facade: Cache) -> None:
backend: Final = azure_blob_facade.cache
assert isinstance(backend, AzureBlobCache)
activate_native(azure_blob_facade)
binding: Final = CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)).resolve()
assert binding.kind == "native"
ping: Final = cast(dict[str, object], await binding.ping())
assert ping["status"] == "success", ping
await binding.async_store(request("async"), {"value": 1})
await binding.async_store({**request("async"), "ttl_seconds": 0.001}, {"value": 2})
time.sleep(0.01)
assert await binding.async_lookup(request("async")) == {"value": 2}
assert await backend.async_get_cache("async") == json.loads(
backend.container_client.download_blob("async").readall()
)
with rebound(azure_blob_facade, "_native_cache", None):
assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="async") == {"value": 2}
await binding.async_store_batch([request("first"), request("second")], [{"value": 3}, {"value": 4}])
assert await binding.async_lookup_batch([request("second"), request("missing"), request("first")]) == {
"values": [{"value": 4}, None, {"value": 3}],
"missing_indices": [1],
}
await binding.async_flush()
assert [blob.name for blob in backend.container_client.list_blobs()] == []
assert await binding.async_lookup(request("async")) is None
async def test_azure_blob_explicit_selection_activates_natively(monkeypatch: pytest.MonkeyPatch) -> None:
account_url: Final = os.environ.get("AZURE_BLOB_CACHE_ACCOUNT_URL")
if account_url is None:
pytest.skip(
"live Azure Blob parity needs AZURE_BLOB_CACHE_ACCOUNT_URL plus DefaultAzureCredential inputs in the environment"
)
facade: Final = activate_native(
Cache(
type=LiteLLMCacheType.AZURE_BLOB,
azure_account_url=account_url,
azure_blob_container=f"litellm-parity-{uuid.uuid4().hex[:12]}",
)
)
backend: Final = facade.cache
assert isinstance(backend, AzureBlobCache)
try:
assert_native_runtime(facade)
kwargs: Final = completion_kwargs("azure")
await facade.async_add_cache({"answer": "azure"}, **kwargs)
assert await facade.async_get_cache(**kwargs) == {"answer": "azure"}
assert backend.get_cache(facade.get_cache_key(**kwargs))["response"] == {"answer": "azure"}
finally:
backend.container_client.delete_container()
await backend.disconnect()

View file

@ -1,115 +0,0 @@
import asyncio
import json
import time
from pathlib import Path
from types import SimpleNamespace
from typing import Final
import diskcache
import pytest
from litellm.caching.caching import Cache
from litellm.caching.disk_cache import DiskCache
from litellm.rust_bridge import _native
from litellm.types.caching import LiteLLMCacheType
from tests.test_litellm_rust.support.cache import CacheTestResolver, activate_native, native_runtime, request
from tests.test_litellm_rust.support.isolation import rebound
pytestmark: Final = pytest.mark.requires_rust_extension
async def test_disk_reads_python_entries_and_python_reads_native_entries(tmp_path: Path) -> None:
disk_cache: Final = DiskCache(disk_cache_dir=str(tmp_path))
response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}}
disk_cache.disk_cache.set(
"sync",
{"timestamp": time.time(), "response": json.dumps(response)},
)
disk_cache.disk_cache.set("async", json.dumps({"timestamp": time.time(), "response": response}))
disk_cache.disk_cache.set("raw", json.dumps(response))
disk_cache.disk_cache.set("invalid", "not a cache entry")
disk_cache.disk_cache.set(
"large",
{"timestamp": time.time(), "response": {"text": "x" * 70_000}},
)
binding: Final = native_runtime(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)))
assert binding.lookup(request("sync")) == response
assert await binding.async_lookup(request("async")) == response
assert binding.lookup(request("raw")) == response
assert await binding.async_lookup(request("invalid")) is None
assert binding.lookup(request("large")) == {"text": "x" * 70_000}
await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response)
stored_response: Final = disk_cache.get_cache("native")
assert isinstance(stored_response, dict)
assert stored_response["response"] == response
stored, expire_time = disk_cache.disk_cache.get("native", expire_time=True)
assert stored is not None
assert time.time() < expire_time <= time.time() + 12.0
await binding.async_store(request("no-ttl"), response)
_, no_expiry = disk_cache.disk_cache.get("no-ttl", expire_time=True)
assert no_expiry is None
async def test_disk_entries_survive_a_fresh_handle_and_expire_on_time(tmp_path: Path) -> None:
first: Final = native_runtime(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)))
await first.async_store(request("persistent"), {"value": "persistent"})
await first.async_store({**request("expiring"), "ttl_seconds": 0.3}, {"value": "expiring"})
fresh: Final = native_runtime(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)))
assert fresh.lookup(request("persistent")) == {"value": "persistent"}
assert fresh.lookup(request("expiring")) == {"value": "expiring"}
await asyncio.sleep(0.4)
assert fresh.lookup(request("expiring")) is None
assert fresh.lookup(request("persistent")) == {"value": "persistent"}
def test_selected_disk_runtime_declines_store_changes(tmp_path: Path) -> None:
facade: Final = activate_native(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)))
selected: Final = CacheTestResolver(SimpleNamespace(cache=facade))
native: Final = selected.resolve()
native.store(request("native"), {"value": "native"})
assert facade.get_cache(cache_key="native") == {"value": "native"}
replacement: Final = diskcache.Cache(str(tmp_path))
try:
with rebound(facade.cache, "disk_cache", replacement):
with pytest.raises(_native.RustBridgeDeclined):
selected.resolve()
assert selected.resolve().kind == "native"
finally:
replacement.close()
class CustomStore(diskcache.Cache):
pass
unsupported: Final = Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path))
store: Final = CustomStore(str(tmp_path))
try:
unsupported.cache.disk_cache = store
with pytest.raises(_native.RustBridgeDeclined, match="built-in diskcache store"):
native_runtime(unsupported)
finally:
store.close()
async def test_disk_native_batch_lookup_and_store_report_partial_hits(tmp_path: Path) -> None:
binding: Final = native_runtime(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)))
requests: Final = [request("hit"), request("miss"), request("disabled")]
requests[2]["controls"] = {
"supported_call_type": True,
"configured": True,
"native_backend": True,
"default_on": True,
"caching": False,
"no_cache": False,
"no_store": False,
"use_cache": False,
}
await binding.async_store_batch(requests, [{"value": 1}, {"value": 2}, {"value": 3}])
partial: Final = await binding.async_lookup_batch(requests)
assert partial == {
"values": [{"value": 1}, {"value": 2}, None],
"missing_indices": [2],
}

View file

@ -1,356 +0,0 @@
import asyncio
import contextvars
import gc
import weakref
from types import SimpleNamespace
from typing import Final, cast
import pytest
import litellm
from litellm.caching.caching import Cache, disable_cache, enable_cache, update_cache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.rust_bridge import _native
from litellm.rust_bridge.response_cache import ResponseCacheRuntime
from litellm.types.caching import LiteLLMCacheType
from tests.test_litellm_rust.support.cache import (
CacheLookup,
CacheTestResolver,
activate_native,
native_runtime,
request,
)
from tests.test_litellm_rust.support.isolation import rebound
pytestmark: Final = pytest.mark.requires_rust_extension
def test_existing_constructor_and_global_are_unchanged() -> None:
facade: Final = Cache(type=LiteLLMCacheType.LOCAL)
assert type(facade.cache) is InMemoryCache
with rebound(litellm, "cache", facade):
resolver: Final = CacheTestResolver(litellm)
assert resolver.resolve().kind == "python_callback"
resolver.resolve().store(None, {"answer": 7}, callback_kwargs={"cache_key": "key"})
assert cast(CacheLookup, facade).get_cache(cache_key="key") == {"answer": 7}
async def test_explicit_selection_constructs_native_runtime_from_public_cache_configuration() -> None:
facade: Final = Cache(type=LiteLLMCacheType.LOCAL)
runtime: Final = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade))
assert isinstance(runtime, ResponseCacheRuntime)
assert runtime.kind == "native"
sync_request: Final = runtime.request(facade, {"cache_key": "sync"})
assert sync_request is not None
runtime.store(sync_request, {"answer": 1})
assert runtime.lookup(sync_request) == {"answer": 1}
assert facade.cache.get_cache("sync") is None
async_request: Final = runtime.request(facade, {"cache_key": "async"})
assert async_request is not None
await runtime.async_store(async_request, {"answer": 2})
assert await runtime.async_lookup(async_request) == {"answer": 2}
assert await facade.cache.async_get_cache("async") is None
requests: Final = (sync_request, async_request)
expected: Final = {
"values": [{"answer": 1}, {"answer": 2}],
"missing_indices": [],
}
assert runtime.lookup_batch(requests) == expected
assert await runtime.async_lookup_batch(requests) == expected
await runtime.async_flush()
assert runtime.lookup(sync_request) is None
assert await runtime.async_lookup(async_request) is None
async def test_inference_resolver_uses_the_configured_native_cache_directly() -> None:
facade: Final = Cache(type=LiteLLMCacheType.LOCAL)
runtime: Final = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade))
assert isinstance(runtime, ResponseCacheRuntime)
facade._native_cache = runtime
selected: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
assert selected.kind == "native"
request: Final = runtime.request(facade, {"cache_key": "inference-native"})
assert request is not None
await selected.async_store(request, {"answer": 42})
assert await selected.async_lookup(request) == {"answer": 42}
assert await runtime.async_lookup(request) == {"answer": 42}
assert facade.cache.get_cache("inference-native") is None
facade._native_cache = None
fallback: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
assert fallback.kind == "python_callback"
await fallback.async_store(None, {"answer": 7}, callback_kwargs={"cache_key": "inference-python"})
assert facade.get_cache(cache_key="inference-python") == {"answer": 7}
assert facade.cache.get_cache("inference-python") is not None
async def test_inference_resolver_declines_a_native_runtime_whose_facade_changed() -> None:
facade: Final = Cache(type=LiteLLMCacheType.LOCAL)
runtime: Final = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade))
assert isinstance(runtime, ResponseCacheRuntime)
facade._native_cache = runtime
stale_request: Final = runtime.request(facade, {"cache_key": "stale-only"})
assert stale_request is not None
await runtime.async_store(stale_request, {"answer": "stale"})
replacement: Final = InMemoryCache()
facade.cache = replacement
with pytest.raises(_native.RustBridgeDeclined):
CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
assert await runtime.async_lookup(stale_request) == {"answer": "stale"}
assert replacement.get_cache("stale-only") is None
assert replacement.get_cache("swapped-backend") is None
def test_existing_global_lifecycle_remains_the_resolver_source_of_truth() -> None:
resolver: Final = CacheTestResolver(litellm)
enable_cache(type=LiteLLMCacheType.LOCAL, ttl=30)
enabled: Final = litellm.cache
assert isinstance(enabled, Cache)
assert enabled.ttl == 30
assert resolver.resolve().kind == "python_callback"
enable_cache(type=LiteLLMCacheType.LOCAL, ttl=60)
assert litellm.cache is enabled
update_cache(type=LiteLLMCacheType.LOCAL, ttl=60)
updated: Final = litellm.cache
assert isinstance(updated, Cache)
assert updated is not enabled
assert updated.ttl == 60
disable_cache()
assert litellm.cache is None
assert resolver.resolve().kind == "disabled"
async def test_native_bindings_survive_replacement_and_capture_writes_before_dispatch() -> None:
namespace: Final = SimpleNamespace(cache=activate_native(Cache(type=LiteLLMCacheType.LOCAL)))
resolver: Final = CacheTestResolver(namespace)
selected: Final = resolver.resolve()
assert selected.kind == "native"
selected.store(request(), {"answer": 1})
assert await selected.async_lookup(request()) == {"answer": 1}
with rebound(namespace, "cache", activate_native(Cache(type=LiteLLMCacheType.LOCAL))):
replacement: Final = resolver.resolve()
await selected.async_store(request(), {"answer": 2})
assert replacement.lookup(request()) is None
assert selected.lookup(request()) == {"answer": 2}
with rebound(namespace, "cache", None):
disabled: Final = resolver.resolve()
assert disabled.kind == "disabled"
assert disabled.lookup(None) is None
await disabled.async_store(None, object())
assert await disabled.async_lookup(None) is None
assert selected.lookup(request()) == {"answer": 2}
async def test_python_callback_preserves_identity_caller_task_context_and_errors() -> None:
context: Final = contextvars.ContextVar("cache_context", default="caller")
caller: Final = asyncio.current_task()
sentinel: Final = object()
failure: Final = RuntimeError("callback failed")
class CustomCache:
async def async_get_cache(self, *, marker: object) -> object:
assert marker is sentinel
assert asyncio.current_task() is caller
context.set("callback")
return marker
async def async_add_cache(self, response: object, *, marker: object) -> None:
assert response is sentinel
assert marker is sentinel
raise failure
namespace: Final = SimpleNamespace(cache=CustomCache())
binding: Final = CacheTestResolver(namespace).resolve()
assert binding.kind == "python_callback"
assert await binding.async_lookup(None, callback_kwargs={"marker": sentinel}) is sentinel
assert context.get() == "callback"
with pytest.raises(RuntimeError) as caught:
await binding.async_store(None, sentinel, callback_kwargs={"marker": sentinel})
assert caught.value is failure
async def test_callback_cancellation_stays_in_the_callers_task() -> None:
entered: Final = asyncio.Event()
finished: Final = asyncio.Event()
class CustomCache:
async def async_get_cache(self) -> None:
entered.set()
try:
await asyncio.Future()
finally:
finished.set()
binding: Final = CacheTestResolver(SimpleNamespace(cache=CustomCache())).resolve()
async def lookup() -> object:
return await binding.async_lookup(None, callback_kwargs={})
task: Final = asyncio.create_task(lookup())
await entered.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert finished.is_set()
@pytest.mark.parametrize("method", ("get_cache", "get_cache_key", "async_get_cache"))
def test_selected_native_runtime_declines_instance_overrides(method: str) -> None:
facade: Final = activate_native(Cache(type=LiteLLMCacheType.LOCAL))
selected: Final = CacheTestResolver(SimpleNamespace(cache=facade))
native: Final = selected.resolve()
native.store(request(), {"source": "native"})
def override(**_kwargs: object) -> None:
return None
with rebound(facade, method, override):
with pytest.raises(_native.RustBridgeDeclined):
selected.resolve()
assert native.lookup(request()) == {"source": "native"}
@pytest.mark.parametrize(("attribute", "value"), (("ttl", 12), ("semantic_cache_scope", "end_user")))
def test_selected_native_runtime_declines_policy_changes(attribute: str, value: object) -> None:
facade: Final = activate_native(Cache(type=LiteLLMCacheType.LOCAL))
selected: Final = CacheTestResolver(SimpleNamespace(cache=facade))
with rebound(facade, attribute, value):
with pytest.raises(_native.RustBridgeDeclined):
selected.resolve()
assert selected.resolve().kind == "native"
def test_resolver_and_callback_cycles_can_be_collected() -> None:
class CustomCache:
pass
def cyclic_reference() -> weakref.ReferenceType[CustomCache]:
callback: Final = CustomCache()
namespace: Final = SimpleNamespace(cache=callback)
binding: Final = CacheTestResolver(namespace).resolve()
setattr(callback, "binding", binding)
return weakref.ref(callback)
reference: Final = cyclic_reference()
gc.collect()
assert reference() is None
def test_invalid_duration_and_request_shape_fail_before_storage() -> None:
binding: Final = CacheTestResolver(
SimpleNamespace(cache=activate_native(Cache(type=LiteLLMCacheType.LOCAL)))
).resolve()
for seconds in (-1.0, float("nan"), float("inf")):
with pytest.raises(ValueError, match="cache durations must be finite and nonnegative"):
binding.store({**request(), "ttl_seconds": seconds}, {"answer": 1})
assert binding.lookup(request()) is None
facade: Final = Cache(type=LiteLLMCacheType.LOCAL)
facade.cache = InMemoryCache(default_ttl=-1)
with pytest.raises(ValueError, match="cache durations must be finite and nonnegative"):
native_runtime(facade)
async def test_memory_size_policy_is_applied_by_the_native_host() -> None:
facade: Final = Cache(type=LiteLLMCacheType.LOCAL)
facade.cache = InMemoryCache(max_size_in_memory=2, max_size_per_item=1)
handle: Final = activate_native(facade)
binding: Final = CacheTestResolver(SimpleNamespace(cache=handle)).resolve()
small: Final = {"answer": "ok"}
binding.store(request("small"), small)
assert await binding.async_lookup(request("small")) == small
await binding.async_store(request("large"), {"answer": "x" * 2048})
assert binding.lookup(request("large")) is None
assert binding.lookup(request("small")) == small
disabled_facade: Final = Cache(type=LiteLLMCacheType.LOCAL)
disabled_facade.cache = InMemoryCache(max_size_in_memory=0)
disabled: Final = native_runtime(disabled_facade)
await disabled.async_store(request(), small)
assert await disabled.async_lookup(request()) is None
async def test_native_batch_lookup_and_store_report_partial_hits() -> None:
binding: Final = CacheTestResolver(
SimpleNamespace(cache=activate_native(Cache(type=LiteLLMCacheType.LOCAL)))
).resolve()
requests: Final = [request("hit"), request("miss"), request("disabled")]
requests[2]["controls"] = {
"supported_call_type": True,
"configured": True,
"native_backend": True,
"default_on": True,
"caching": False,
"no_cache": False,
"no_store": False,
"use_cache": False,
}
await binding.async_store_batch(requests, [{"value": 1}, {"value": 2}, {"value": 3}])
partial: Final = await binding.async_lookup_batch(requests)
assert partial == {
"values": [{"value": 1}, {"value": 2}, None],
"missing_indices": [2],
}
async def test_python_batch_callbacks_use_the_builtin_cache_api() -> None:
result: Final = object()
marker: Final = object()
class CustomCache(Cache):
def get_cache(self, dynamic_cache_object: object = None, **kwargs: object) -> object:
return ("sync", kwargs)
async def async_get_cache(self, dynamic_cache_object: object = None, **kwargs: object) -> object:
return ("async", kwargs)
async def async_add_cache_pipeline(
self, result: object, dynamic_cache_object: object = None, **kwargs: object
) -> object:
return result, kwargs
binding: Final = CacheTestResolver(SimpleNamespace(cache=CustomCache(type=LiteLLMCacheType.LOCAL))).resolve()
assert binding.kind == "python_callback"
requests: Final = [request("first"), request("second")]
kwargs: Final = [{"cache_key": "first"}, {"cache_key": "second"}]
assert binding.lookup_batch(requests, callback_kwargs=kwargs) == [("sync", kwargs[0]), ("sync", kwargs[1])]
assert await binding.async_lookup_batch(requests, callback_kwargs=kwargs) == [
("async", kwargs[0]),
("async", kwargs[1]),
]
with pytest.raises(ValueError, match="equal lengths"):
binding.lookup_batch(requests, callback_kwargs=kwargs[:1])
with pytest.raises(TypeError, match="callback_result"):
await binding.async_store_batch(requests, [1, 2], callback_kwargs={"marker": marker})
stored: Final = cast(
tuple[object, dict[str, object]],
await binding.async_store_batch(requests, [1, 2], callback_result=result, callback_kwargs={"marker": marker}),
)
assert stored[0] is result
assert stored[1] == {"marker": marker}
async def test_unmodified_builtin_cache_callbacks_can_ping_and_flush() -> None:
async def ping() -> str:
return "pong"
cache: Final = Cache(type=LiteLLMCacheType.LOCAL)
cache.cache.set_cache("key", "value")
binding: Final = CacheTestResolver(SimpleNamespace(cache=cache)).resolve()
assert binding.kind == "python_callback"
setattr(cache.cache, "ping", ping)
assert await binding.ping() == "pong"
await binding.async_flush()
assert cache.cache.get_cache("key") is None

View file

@ -1,46 +0,0 @@
from types import SimpleNamespace
from typing import Final
import pytest
from litellm.caching.caching import Cache
from litellm.rust_bridge import _native
from litellm.types.caching import LiteLLMCacheType
from tests.test_litellm_rust.support.cache import CacheTestResolver, activate_native, native_runtime
from tests.test_litellm_rust.support.isolation import rebound
pytestmark: Final = pytest.mark.requires_rust_extension
@pytest.mark.parametrize(
("attribute", "replacement"),
(("bucket_name", "other"), ("key_prefix", "other/"), ("path_service_account", "other.json")),
)
def test_selected_gcs_runtime_declines_backend_configuration_changes(
monkeypatch: pytest.MonkeyPatch, attribute: str, replacement: str
) -> None:
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
facade: Final = activate_native(Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/"))
selected: Final = CacheTestResolver(SimpleNamespace(cache=facade))
assert selected.resolve().kind == "native"
with rebound(facade.cache, attribute, replacement):
with pytest.raises(_native.RustBridgeDeclined):
selected.resolve()
assert selected.resolve().kind == "native"
async def test_gcs_runtime_flush_is_a_no_op_and_ping_is_not_implemented(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
runtime: Final = native_runtime(Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket"))
await runtime.async_flush()
with pytest.raises(NotImplementedError):
await runtime.ping()
def test_gcs_runtime_declines_missing_bucket_configuration(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False)
monkeypatch.delenv("GCS_BUCKET_NAME", raising=False)
with pytest.raises(_native.RustBridgeDeclined, match="requires a configured bucket name"):
native_runtime(Cache(type=LiteLLMCacheType.GCS))

View file

@ -1,254 +0,0 @@
import hashlib
import http.server
import json
import math
import os
import threading
import time
from collections.abc import Generator
from types import SimpleNamespace
from typing import Final
from uuid import uuid4
import pytest
from litellm.caching.caching import Cache
from litellm.rust_bridge import _native
from litellm.types.caching import LiteLLMCacheType
from tests.test_litellm_rust.support.cache import (
CacheTestResolver,
activate_native,
assert_native_runtime,
native_runtime,
request,
)
pytestmark: Final = pytest.mark.requires_rust_extension
def qdrant_request(
key: str,
messages: list[dict[str, object]],
**kwargs: object,
) -> dict[str, object]:
return {**request(key), "messages": messages, **kwargs}
def embedding_vector(text: str) -> list[float]:
raw: Final = hashlib.sha256(text.encode()).digest()[:8]
values: Final = [byte / 127.5 - 1 for byte in raw]
norm: Final = math.sqrt(sum(value * value for value in values))
return [value / norm for value in values]
@pytest.fixture
def qdrant_url() -> str:
value: Final[str | None] = os.environ.get("QDRANT_URL")
if not value:
pytest.skip("QDRANT_URL is required for Qdrant semantic cache tests")
return value.rstrip("/")
@pytest.fixture
def fake_embedding_endpoint(monkeypatch: pytest.MonkeyPatch) -> Generator[str]:
class EmbeddingHandler(http.server.BaseHTTPRequestHandler):
def do_POST(self) -> None:
length: Final = int(self.headers["Content-Length"])
body: Final = json.loads(self.rfile.read(length))
text: Final = body["input"]
response: Final = {
"object": "list",
"data": [
{
"object": "embedding",
"index": 0,
"embedding": embedding_vector(text),
}
],
"model": body["model"],
"usage": {"prompt_tokens": 1, "total_tokens": 1},
}
encoded: Final = json.dumps(response).encode()
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(encoded)))
self.end_headers()
self.wfile.write(encoded)
def log_message(self, *_args: object) -> None:
return
server: Final = http.server.ThreadingHTTPServer(("127.0.0.1", 0), EmbeddingHandler)
worker: Final = threading.Thread(target=server.serve_forever, daemon=True)
worker.start()
monkeypatch.setenv("OPENAI_API_BASE", f"http://127.0.0.1:{server.server_address[1]}")
monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
try:
yield f"http://127.0.0.1:{server.server_address[1]}"
finally:
server.shutdown()
server.server_close()
worker.join(timeout=5)
def qdrant_facade(qdrant_url: str, collection_name: str) -> Cache:
return Cache(
type=LiteLLMCacheType.QDRANT_SEMANTIC,
qdrant_api_base=qdrant_url,
qdrant_collection_name=collection_name,
similarity_threshold=0.99,
qdrant_semantic_cache_embedding_model="text-embedding-3-small",
qdrant_semantic_cache_vector_size=8,
)
def test_qdrant_semantic_facade_binds_native_and_shares_entries(qdrant_url: str, fake_embedding_endpoint: str) -> None:
del fake_embedding_endpoint
messages: Final = [{"role": "user", "content": "shared prompt"}]
collection: Final = f"cache_{uuid4().hex}"
facade: Final = qdrant_facade(qdrant_url, collection)
facade.cache.set_cache(
"python-key",
{"timestamp": time.time(), "response": json.dumps({"id": "py"})},
messages=messages,
)
activate_native(facade)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
assert binding.kind == "native"
assert binding.lookup(qdrant_request("python-key", messages)) == {"id": "py"}
binding.store(qdrant_request("native-key", messages), {"id": "native"})
python_value: Final = facade.cache.get_cache("native-key", messages=messages)
assert isinstance(python_value, dict)
assert python_value["response"] == {"id": "native"}
unrelated: Final = [{"role": "user", "content": "unrelated prompt"}]
assert binding.lookup(qdrant_request("native-key", unrelated)) is None
assert facade.cache.get_cache("native-key", messages=unrelated) is None
assert binding.lookup(qdrant_request("different-key", messages)) is None
assert facade.cache.get_cache("different-key", messages=messages) is None
async def test_qdrant_semantic_async_parity(qdrant_url: str, fake_embedding_endpoint: str) -> None:
del fake_embedding_endpoint
messages: Final = [{"role": "user", "content": "async prompt"}]
collection: Final = f"cache_{uuid4().hex}"
facade: Final = qdrant_facade(qdrant_url, collection)
activate_native(facade)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
await facade.cache.async_set_cache(
"python-key",
{"timestamp": time.time(), "response": json.dumps({"id": "py"})},
messages=messages,
)
assert await binding.async_lookup(qdrant_request("python-key", messages)) == {"id": "py"}
await binding.async_store(qdrant_request("native-key", messages), {"id": "native"})
python_value: Final = await facade.cache.async_get_cache("native-key", messages=messages)
assert isinstance(python_value, dict)
assert python_value["response"] == {"id": "native"}
async def test_qdrant_semantic_async_store_batch_shares_entries(qdrant_url: str, fake_embedding_endpoint: str) -> None:
del fake_embedding_endpoint
collection: Final = f"cache_{uuid4().hex}"
facade: Final = qdrant_facade(qdrant_url, collection)
activate_native(facade)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
entries: Final = [
qdrant_request("batch-one", [{"role": "user", "content": "first batch prompt"}]),
qdrant_request("batch-two", [{"role": "user", "content": "second batch prompt"}]),
]
await binding.async_store_batch(entries, [{"id": "one"}, {"id": "two"}])
assert binding.lookup(entries[0]) == {"id": "one"}
assert binding.lookup(entries[1]) == {"id": "two"}
assert (await facade.cache.async_get_cache("batch-one", messages=entries[0]["messages"]))["response"] == {
"id": "one"
}
assert (await facade.cache.async_get_cache("batch-two", messages=entries[1]["messages"]))["response"] == {
"id": "two"
}
async def test_qdrant_semantic_malformed_entries_and_unsupported_operations(
qdrant_url: str, fake_embedding_endpoint: str
) -> None:
del fake_embedding_endpoint
messages: Final = [{"role": "user", "content": "malformed prompt"}]
collection: Final = f"cache_{uuid4().hex}"
facade: Final = qdrant_facade(qdrant_url, collection)
activate_native(facade)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
key: Final = "malformed-key"
response: Final = {
"points": [
{
"id": str(uuid4()),
"vector": embedding_vector("malformed prompt"),
"payload": {
"litellm_cache_key": key,
"text": "malformed prompt",
"response": "not json",
},
}
]
}
facade.cache.sync_client.put(
url=f"{qdrant_url}/collections/{collection}/points",
headers=facade.cache.headers,
json=response,
)
assert binding.lookup(qdrant_request(key, messages)) is None
with pytest.raises(RuntimeError, match="operation is not supported"):
binding.lookup_batch([qdrant_request(key, messages)])
with pytest.raises(RuntimeError, match="operation is not supported"):
await binding.async_flush()
with pytest.raises(RuntimeError, match="operation is not supported"):
await binding.ping()
def test_qdrant_semantic_ignores_request_expiry(qdrant_url: str, fake_embedding_endpoint: str) -> None:
del fake_embedding_endpoint
messages: Final = [{"role": "user", "content": "persistent prompt"}]
collection: Final = f"cache_{uuid4().hex}"
facade: Final = qdrant_facade(qdrant_url, collection)
activate_native(facade)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
binding.store(qdrant_request("persistent-key", messages, ttl_seconds=1.0), {"id": "persistent"})
time.sleep(1.2)
assert binding.lookup(qdrant_request("persistent-key", messages)) == {"id": "persistent"}
python_value: Final = facade.cache.get_cache("persistent-key", messages=messages)
assert isinstance(python_value, dict)
assert python_value["response"] == {"id": "persistent"}
def test_qdrant_runtime_declines_mutation_and_unsupported_configuration(
qdrant_url: str, fake_embedding_endpoint: str
) -> None:
del fake_embedding_endpoint
collection: Final = f"cache_{uuid4().hex}"
facade: Final = qdrant_facade(qdrant_url, collection)
activate_native(facade)
facade.cache.qdrant_api_key = "rotated"
with pytest.raises(_native.RustBridgeDeclined):
CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
facade.cache.similarity_threshold = 0.5
with pytest.raises(_native.RustBridgeDeclined):
CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
unsupported: Final = qdrant_facade(qdrant_url, f"cache_{uuid4().hex}")
unsupported.cache.embedding_max_input_tokens = 100
with pytest.raises(_native.RustBridgeDeclined, match="requires Python"):
native_runtime(unsupported)
unsupported.cache.embedding_max_input_tokens = None
unsupported.cache.qdrant_api_base = "http://127.0.0.1:7777"
with pytest.raises(_native.RustBridgeDeclined, match="gRPC"):
native_runtime(unsupported)
def test_qdrant_semantic_explicit_selection_activates_natively(
qdrant_url: str, fake_embedding_endpoint: str, monkeypatch: pytest.MonkeyPatch
) -> None:
del fake_embedding_endpoint
facade: Final = activate_native(qdrant_facade(qdrant_url, f"cache_{uuid4().hex}"))
assert_native_runtime(facade)
kwargs: Final = {"model": "gpt-4o", "messages": [{"role": "user", "content": "qdrant activation"}]}
facade.add_cache({"answer": "qdrant"}, **kwargs)
assert facade.get_cache(**kwargs) == {"answer": "qdrant"}

View file

@ -1,212 +0,0 @@
import json
import os
import time
from types import SimpleNamespace
from typing import Final
from urllib.parse import urlparse
import pytest
import redis
import litellm
from litellm.caching.caching import Cache
from litellm.caching.redis_cluster_cache import RedisClusterCache
from litellm.rust_bridge import _native
from litellm.types.caching import LiteLLMCacheType
from tests.test_litellm_rust.support.cache import (
CacheTestResolver,
activate_native,
assert_native_runtime,
completion_kwargs,
native_runtime,
request,
)
from tests.test_litellm_rust.support.isolation import rebound
pytestmark: Final = pytest.mark.requires_rust_extension
@pytest.fixture
def cluster_nodes() -> tuple[tuple[str, int], ...]:
configured: Final = os.environ.get("LITELLM_TEST_REDIS_CLUSTER_NODES")
if not configured:
pytest.skip("LITELLM_TEST_REDIS_CLUSTER_NODES is not set")
return tuple((host, int(port)) for host, _, port in (node.partition(":") for node in configured.split(",")))
async def test_redis_reads_python_sync_and_async_entries_and_writes_without_hidden_prefix(redis_url: str) -> None:
client: Final = redis.Redis.from_url(redis_url)
binding: Final = native_runtime(redis_facade(redis_url, namespace="team"))
response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None}
envelope: Final = {"timestamp": time.time(), "response": json.dumps(response)}
client.set("team:sync", str(envelope))
client.set("team:async", json.dumps({"timestamp": time.time(), "response": response}))
client.set("team:raw", json.dumps(response))
client.set("team:invalid", "not a cache entry")
assert binding.lookup(request("sync")) == response
assert await binding.async_lookup(request("team:async")) == response
assert binding.lookup(request("raw")) == response
assert await binding.async_lookup(request("invalid")) is None
await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response)
stored: Final = client.get("team:native")
assert isinstance(stored, bytes)
assert json.loads(stored)["response"] == response
assert 0 < client.ttl("team:native") <= 12
assert client.get("litellm-cache:team:native") is None
assert client.get("team:team:async") is None
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):
facade: Final = Cache(
type=LiteLLMCacheType.REDIS,
host=parsed.hostname,
port=str(parsed.port),
redis_flush_size=2,
)
activate_native(facade)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
client: Final = redis.Redis.from_url(redis_url)
with rebound(facade.cache, "redis_kwargs", {**facade.cache.redis_kwargs, "ssl": True}):
with pytest.raises(_native.RustBridgeDeclined):
CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
pool: Final = facade.cache.redis_client.connection_pool
with rebound(pool, "connection_kwargs", {**pool.connection_kwargs, "db": 1}):
with pytest.raises(_native.RustBridgeDeclined):
CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
await binding.async_store(request("first"), {"value": 1})
assert client.get("first") is None
await binding.async_store(request("second"), {"value": 2})
assert client.get("first") is not None
assert client.get("second") is not None
await facade.cache.disconnect()
client.close()
async def test_redis_cluster_facade_serves_multi_slot_batches_and_scoped_flush_natively(
cluster_nodes: tuple[tuple[str, int], ...],
) -> None:
startup_nodes: Final = [{"host": host, "port": port} for host, port in cluster_nodes]
with rebound(litellm, "default_redis_ttl", 60):
facade: Final = Cache(type=LiteLLMCacheType.REDIS, redis_startup_nodes=startup_nodes, namespace="parity")
assert type(facade.cache) is RedisClusterCache
activate_native(facade)
resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade))
assert resolver.resolve().kind == "native"
manager: Final = facade.cache.redis_client.nodes_manager
with rebound(manager, "connection_kwargs", {**manager.connection_kwargs, "db": 1}):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
with rebound(facade.cache, "redis_kwargs", {**facade.cache.redis_kwargs, "startup_nodes": startup_nodes[:1]}):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
binding: Final = resolver.resolve()
assert binding.kind == "native"
client: Final = redis.RedisCluster(startup_nodes=[redis.cluster.ClusterNode(*node) for node in cluster_nodes])
keys: Final = tuple(f"slot-{index}" for index in range(12))
slots: Final = {client.keyslot(f"parity:{key}") for key in keys}
assert len(slots) > 1, slots
requests: Final = [request(key) for key in keys]
values: Final = [{"index": index} for index in range(len(keys))]
await binding.async_store_batch(requests, values)
client.set("parity:slot-3", "not a cache entry")
client.set("parity:slot-7", json.dumps({"timestamp": time.time(), "response": {"index": 7, "python": True}}))
batch: Final = await binding.async_lookup_batch(requests)
assert batch == {
"values": [
None if index == 3 else {"index": 7, "python": True} if index == 7 else value
for index, value in enumerate(values)
],
"missing_indices": [3],
}
assert facade.cache.get_cache("parity:slot-0")["response"] == {"index": 0}
assert (await facade.cache.async_get_cache("parity:slot-11"))["response"] == {"index": 11}
assert facade.cache.redis_client.mget_nonatomic([f"parity:{key}" for key in keys[:2]]) == [
client.get("parity:slot-0"),
client.get("parity:slot-1"),
]
await binding.async_store({**request("pinned"), "ttl_seconds": 12.0}, {"pinned": True})
assert 0 < client.ttl("parity:pinned") <= 12
client.set("unscoped", "stays")
await binding.async_flush()
remaining: Final = tuple(
sorted(key for node in client.get_primaries() for key in client.keys("parity:*", target_nodes=node))
)
assert remaining == (), remaining
assert client.get("unscoped") == b"stays"
client.delete("unscoped")
client.close()
facade.cache.redis_client.close()
def redis_facade(redis_url: str, **settings: object) -> Cache:
parsed: Final = urlparse(redis_url)
return Cache(type=LiteLLMCacheType.REDIS, host=parsed.hostname, port=str(parsed.port), **settings)
@pytest.mark.parametrize(
("settings", "message"),
[
pytest.param({"max_connections": 10}, "max_connections requires Python", id="pool-size"),
pytest.param({"socket_timeout": 1.0}, "socket_timeout and socket_connect_timeout", id="socket-timeout"),
pytest.param(
{"socket_connect_timeout": 1.0}, "socket_timeout and socket_connect_timeout", id="connect-timeout"
),
pytest.param({"socket_keepalive": True}, "does not support socket_keepalive", id="keepalive"),
pytest.param({"health_check_interval": 5}, "does not support health_check_interval", id="health-check"),
pytest.param({"client_name": "litellm"}, "does not support client_name", id="client-name"),
pytest.param({"ssl": True}, "ssl_check_hostname=false require Python", id="tls-default-hostname-check"),
pytest.param({"ssl": True, "ssl_cert_reqs": "none"}, "ssl_cert_reqs=none", id="tls-without-verification"),
pytest.param(
{"ssl": True, "ssl_check_hostname": True, "ssl_ca_certs": "/ca.pem"},
"does not support ssl_ca_certs",
id="tls-custom-ca",
),
pytest.param(
{"ssl": True, "ssl_check_hostname": True, "ssl_certfile": "/client.pem", "ssl_keyfile": "/client.key"},
"does not support ssl_ca_certs, ssl_ca_data, ssl_certfile or ssl_keyfile",
id="tls-client-certificate",
),
],
)
def test_redis_settings_the_native_client_cannot_honor_decline(
redis_url: str, monkeypatch: pytest.MonkeyPatch, settings: dict[str, object], message: str
) -> None:
with pytest.raises(_native.RustBridgeDeclined, match=f"native Redis.*{message}"):
activate_native(redis_facade(redis_url, **settings))
def test_redis_verified_tls_activates_natively(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None:
assert_native_runtime(activate_native(redis_facade(redis_url, ssl=True, ssl_check_hostname=True)))
async def test_redis_flush_size_buffers_native_facade_writes(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None:
facade: Final = activate_native(redis_facade(redis_url, redis_flush_size=2, namespace="team"))
assert_native_runtime(facade)
client: Final = redis.Redis.from_url(redis_url)
first: Final = completion_kwargs("first")
await facade.async_add_cache({"value": 1}, **first)
first_key: Final = facade.get_cache_key(**first)
assert first_key.startswith("team:")
assert client.get(first_key) is None
second: Final = completion_kwargs("second")
await facade.async_add_cache({"value": 2}, **second)
assert client.get(first_key) is not None
assert client.get(facade.get_cache_key(**second)) is not None
client.close()
def test_legacy_constructor_accepts_python_only_settings(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None:
assert redis_facade(redis_url, socket_timeout=1.0)._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor

View file

@ -1,560 +0,0 @@
import asyncio
import contextvars
import hashlib
import json
import math
import os
from collections.abc import Callable, Generator
from contextlib import ExitStack
from types import SimpleNamespace
from typing import Final, cast
from uuid import uuid4
import pytest
import redis
import litellm
from litellm.caching.caching import Cache
from litellm.caching.redis_semantic_cache import RedisSemanticCache
from litellm.rust_bridge import _native
from litellm.types.caching import LiteLLMCacheType
from litellm.types.llms.custom_llm import CustomLLMItem
from litellm.types.utils import EmbeddingResponse
from tests.test_litellm_rust.support.cache import (
CacheTestResolver,
activate_native,
assert_native_runtime,
request,
)
from tests.test_litellm_rust.support.isolation import rebound
pytestmark: Final = pytest.mark.requires_rust_extension
PARAPHRASE_MARKER: Final = " (paraphrase)"
SEMANTIC_EMBEDDING_MODEL: Final = "semantic-test/deterministic"
SEMANTIC_INDEX_PREFIX: Final = "litellm_test_semantic_"
SEMANTIC_CONTEXT: Final = contextvars.ContextVar("semantic_test_context", default="unset")
def _normalized(vector: list[float]) -> list[float]:
norm: Final = math.sqrt(sum(component * component for component in vector))
return [component / norm for component in vector]
def _base_embedding(prompt: str) -> list[float]:
digest: Final = hashlib.sha256(prompt.encode("utf-8")).digest()
return _normalized([float(digest[index] + 1) for index in range(8)])
def _semantic_embedding(prompt: str) -> list[float]:
if PARAPHRASE_MARKER not in prompt:
return _base_embedding(prompt)
base: Final = _base_embedding(prompt.replace(PARAPHRASE_MARKER, "").strip())
pivot: Final = min(range(8), key=lambda index: abs(base[index]))
direction: Final = _normalized(
[(1.0 - base[pivot] * base[pivot]) if index == pivot else -base[index] * base[pivot] for index in range(8)]
)
# Rotating an orthogonal unit direction by 0.329 produces ~0.05 cosine distance
return _normalized([base[index] + 0.329 * direction[index] for index in range(8)])
class DeterministicEmbedding(litellm.CustomLLM):
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
self.async_calls: list[dict[str, object]] = []
self.entered = asyncio.Event()
self.gate: asyncio.Event | None = None
def _respond(
self,
model: str,
input: object,
model_response: EmbeddingResponse,
) -> EmbeddingResponse:
texts: Final = cast(list[object], input if isinstance(input, list) else [input])
self.calls.append({"model": model, "input": texts})
model_response.model = model
model_response.data = [
{"object": "embedding", "index": index, "embedding": _semantic_embedding(str(text))}
for index, text in enumerate(texts)
]
return model_response
def embedding(
self,
model: str,
input: list[object],
model_response: EmbeddingResponse,
print_verbose: Callable[..., object],
logging_obj: object,
optional_params: dict[str, object],
api_key: object = None,
api_base: object = None,
timeout: object = None,
litellm_params: object = None,
) -> EmbeddingResponse:
return self._respond(model, input, model_response)
async def aembedding(
self,
model: str,
input: list[object],
model_response: EmbeddingResponse,
print_verbose: Callable[..., object],
logging_obj: object,
optional_params: dict[str, object],
api_key: object = None,
api_base: object = None,
timeout: object = None,
litellm_params: object = None,
) -> EmbeddingResponse:
texts: Final = cast(list[object], input if isinstance(input, list) else [input])
self.async_calls.append(
{
"model": model,
"input": texts,
"task": asyncio.current_task(),
"context": SEMANTIC_CONTEXT.get(),
}
)
SEMANTIC_CONTEXT.set("written-in-aembedding")
self.entered.set()
if self.gate is not None:
await self.gate.wait()
return self._respond(model, input, model_response)
@pytest.fixture
def semantic_embedding() -> Generator[DeterministicEmbedding]:
handler: Final = DeterministicEmbedding()
with ExitStack() as stack:
stack.enter_context(
rebound(
litellm,
"custom_provider_map",
[
*litellm.custom_provider_map,
cast(
CustomLLMItem,
{"provider": "semantic-test", "custom_handler": handler},
),
],
)
)
stack.enter_context(
rebound(
litellm,
"_custom_providers", # pyright: ignore[reportPrivateUsage] # no public provider-registration hook
[*litellm._custom_providers, "semantic-test"], # pyright: ignore[reportPrivateUsage] # no public provider-registration hook
)
)
stack.enter_context(rebound(litellm, "provider_list", [*litellm.provider_list, "semantic-test"]))
yield handler
@pytest.fixture
def redis_stack() -> Generator[tuple[str, str]]:
url: Final = os.environ.get("LITELLM_REDIS_STACK_URL")
if url is None:
pytest.skip("LITELLM_REDIS_STACK_URL is not set")
index: Final = f"{SEMANTIC_INDEX_PREFIX}{uuid4().hex}"
yield url, index
client: Final = redis.Redis.from_url(url)
try:
client.execute_command("FT.DROPINDEX", index, "DD") # pyright: ignore[reportUnknownMemberType] # redis-py leaves execute_command partially unknown
except redis.RedisError:
pass
client.close()
def semantic_request(key: str, prompt: str, **extra: object) -> dict[str, object]:
return {
"key": {"preset": key},
"messages": [{"role": "user", "content": prompt}],
**extra,
}
def semantic_messages(prompt: str) -> list[dict[str, object]]:
return [{"role": "user", "content": prompt}]
def semantic_entry_id(prompt: str, tag: str) -> str:
return hashlib.sha256(f"{prompt}litellm_cache_key{tag}".encode()).hexdigest()
def semantic_facade(url: str, index: str, *, similarity_threshold: float = 0.8) -> Cache:
facade: Final = Cache(
type=LiteLLMCacheType.REDIS_SEMANTIC,
redis_url=url,
similarity_threshold=similarity_threshold,
redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL,
redis_semantic_cache_index_name=index,
)
activate_native(facade)
return facade
def test_redis_semantic_constructor_identity_and_provenance(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
backend: Final = cast(RedisSemanticCache, facade.cache)
assert backend.__class__.__module__ == "litellm.caching.redis_semantic_cache"
assert type(backend) is RedisSemanticCache
assert backend._redis_url == url # pyright: ignore[reportPrivateUsage] # provenance check needs the projected config
assert backend._index_name == index # pyright: ignore[reportPrivateUsage] # provenance check needs the projected config
assert backend.similarity_threshold == 0.8
assert backend.embedding_model == SEMANTIC_EMBEDDING_MODEL
assert_native_runtime(facade)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
assert binding.kind == "native"
def test_redis_semantic_native_and_python_sync_entries_share_one_layout(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
client: Final = redis.Redis.from_url(url)
response: Final = {"choices": [{"text": "paris"}], "usage": {"total_tokens": 2}}
binding.store(semantic_request("geo", "what is the capital of france"), response)
native_hash_key: Final = f"{index}:{semantic_entry_id('what is the capital of france', 'geo')}"
stored: Final = client.hgetall(native_hash_key)
assert set(stored) == {
b"entry_id",
b"prompt",
b"response",
b"prompt_vector",
b"inserted_at",
b"updated_at",
b"litellm_cache_key",
}, stored
assert stored[b"entry_id"].decode() == native_hash_key.split(":", 1)[1]
assert stored[b"prompt"] == b"what is the capital of france"
assert stored[b"litellm_cache_key"] == b"geo"
assert len(stored[b"prompt_vector"]) == 32
decoded: Final = cast(dict[str, object], json.loads(stored[b"response"]))
assert decoded["response"] == response
assert (
cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class
"geo", messages=semantic_messages("what is the capital of france")
)
== decoded
)
assert semantic_embedding.calls == [
{"model": "deterministic", "input": ["what is the capital of france"]},
{"model": "deterministic", "input": ["what is the capital of france"]},
{"model": "deterministic", "input": ["dimension test"]},
]
cast(RedisSemanticCache, facade.cache).set_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class
"math",
json.dumps({"timestamp": 1700000000.0, "response": {"answer": 42}}),
messages=semantic_messages("what is 6 times 7"),
)
python_hash_key: Final = f"{index}:{semantic_entry_id('what is 6 times 7', 'math')}"
assert json.loads(cast(bytes, client.hget(python_hash_key, "response"))) == {
"timestamp": 1700000000.0,
"response": {"answer": 42},
}
assert binding.lookup(semantic_request("math", "what is 6 times 7")) == {"answer": 42}
client.close()
async def test_redis_semantic_async_paths_and_store_batch_share_one_layout(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
client: Final = redis.Redis.from_url(url)
await binding.async_store(semantic_request("async", "name a primary color"), {"answer": "blue"})
hash_key: Final = f"{index}:{semantic_entry_id('name a primary color', 'async')}"
decoded: Final = cast(dict[str, object], json.loads(cast(bytes, client.hget(hash_key, "response"))))
python_read: Final = await cast(RedisSemanticCache, facade.cache).async_get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class
"async", messages=semantic_messages("name a primary color")
)
assert python_read == decoded
await binding.async_store_batch(
[
semantic_request("batch-one", "first batch prompt"),
semantic_request("batch-two", "second batch prompt"),
],
[{"answer": 1}, {"answer": 2}],
)
expected: Final = {
key: json.loads(cast(bytes, client.hget(f"{index}:{semantic_entry_id(prompt, key)}", "response")))
for key, prompt in (
("batch-one", "first batch prompt"),
("batch-two", "second batch prompt"),
)
}
for key, prompt in (
("batch-one", "first batch prompt"),
("batch-two", "second batch prompt"),
):
assert (
cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class
key, messages=semantic_messages(prompt)
)
== expected[key]
), key
cast(RedisSemanticCache, facade.cache).set_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class
"async-python",
json.dumps({"timestamp": 1700000000.0, "response": {"answer": "python"}}),
messages=semantic_messages("python written prompt"),
)
assert await binding.async_lookup(semantic_request("async-python", "python written prompt")) == {"answer": "python"}
client.close()
async def test_native_semantic_async_embedding_runs_inline_in_the_callers_task(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
assert binding.kind == "native"
caller: Final = asyncio.current_task()
SEMANTIC_CONTEXT.set("caller-sentinel")
response: Final = {"choices": [{"text": "paris"}]}
await binding.async_store(semantic_request("inline", "what is the capital of france"), response)
assert (
await binding.async_lookup(semantic_request("inline", f"what is the capital of france{PARAPHRASE_MARKER}"))
== response
)
assert await binding.async_lookup(semantic_request("inline", "python written prompt")) is None
assert SEMANTIC_CONTEXT.get() == "written-in-aembedding"
assert semantic_embedding.async_calls == [
{
"model": "deterministic",
"input": ["what is the capital of france"],
"task": caller,
"context": "caller-sentinel",
},
{
"model": "deterministic",
"input": [f"what is the capital of france{PARAPHRASE_MARKER}"],
"task": caller,
"context": "written-in-aembedding",
},
{
"model": "deterministic",
"input": ["python written prompt"],
"task": caller,
"context": "written-in-aembedding",
},
], semantic_embedding.async_calls
async def test_native_semantic_cancellation_during_embedding_skips_the_backend(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
assert binding.kind == "native"
semantic_embedding.gate = asyncio.Event()
async def lookup() -> object:
return await binding.async_lookup(semantic_request("cancel", "cancelled prompt"))
task: Final = asyncio.create_task(lookup())
await semantic_embedding.entered.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
semantic_embedding.gate.set()
assert len(semantic_embedding.async_calls) == 1
assert (
await cast(RedisSemanticCache, facade.cache).async_get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class
"cancel", messages=semantic_messages("cancelled prompt")
)
is None
)
def test_redis_semantic_similarity_tag_and_threshold_boundaries(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
binding.store(semantic_request("sim", "tell me a joke"), {"answer": "haha"})
paraphrase: Final = f"tell me a joke{PARAPHRASE_MARKER}"
assert binding.lookup(semantic_request("sim", paraphrase)) == {"answer": "haha"}
assert binding.lookup(semantic_request("sim", "an unrelated question about spreadsheets")) is None
assert binding.lookup(semantic_request("other-key", "tell me a joke")) is None
strict: Final = semantic_facade(url, index, similarity_threshold=0.99)
strict_binding: Final = CacheTestResolver(SimpleNamespace(cache=strict)).resolve()
assert strict_binding.lookup(semantic_request("sim", paraphrase)) is None
assert strict_binding.lookup(semantic_request("sim", "tell me a joke")) == {"answer": "haha"}
def test_redis_semantic_ttl_is_written_only_when_requested(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
client: Final = redis.Redis.from_url(url)
binding.store({**semantic_request("ttl", "ttl prompt"), "ttl_seconds": 12.0}, {"answer": 1})
expiring: Final = f"{index}:{semantic_entry_id('ttl prompt', 'ttl')}"
assert 0 < client.ttl(expiring) <= 12
binding.store(semantic_request("ttl-none", "untimed prompt"), {"answer": 2})
persistent: Final = f"{index}:{semantic_entry_id('untimed prompt', 'ttl-none')}"
assert client.ttl(persistent) == -1
binding.store(
{**semantic_request("ttl-fraction", "fractional prompt"), "ttl_seconds": 1.5},
{"answer": 3},
)
fractional: Final = f"{index}:{semantic_entry_id('fractional prompt', 'ttl-fraction')}"
assert client.ttl(fractional) == 2
client.close()
def test_redis_semantic_malformed_response_is_a_miss_for_both_readers(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
client: Final = redis.Redis.from_url(url)
binding.store(semantic_request("bad", "corrupt me"), {"answer": 1})
hash_key: Final = f"{index}:{semantic_entry_id('corrupt me', 'bad')}"
client.hset(hash_key, "response", b"{not json")
assert binding.lookup(semantic_request("bad", "corrupt me")) is None
assert (
cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class
"bad", messages=semantic_messages("corrupt me")
)
is None
)
client.close()
async def test_redis_semantic_unsupported_operations_raise_not_implemented(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
with pytest.raises(NotImplementedError):
binding.lookup_batch([semantic_request("batch", "prompt one")])
with pytest.raises(NotImplementedError):
await binding.async_lookup_batch([semantic_request("batch", "prompt one")])
with pytest.raises(NotImplementedError):
await binding.async_flush()
with pytest.raises(NotImplementedError):
await binding.ping()
def test_redis_semantic_requests_without_prompt_are_noops(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
client: Final = redis.Redis.from_url(url)
binding.store(request("plain"), {"answer": 1})
assert binding.lookup(request("plain")) is None
assert semantic_embedding.calls == []
assert client.keys(f"{index}:*") == []
client.close()
def test_redis_semantic_scope_overrides_the_tag_and_isolates_entries(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve()
client: Final = redis.Redis.from_url(url)
scoped: Final = {**semantic_request("scoped", "scoped prompt"), "scope": "team-a"}
binding.store(scoped, {"answer": "kept"})
hash_key: Final = f"{index}:{semantic_entry_id('scoped prompt', 'team-a')}"
assert client.hget(hash_key, "litellm_cache_key") == b"team-a"
assert binding.lookup(scoped) == {"answer": "kept"}
assert binding.lookup(semantic_request("scoped", "scoped prompt")) is None
assert binding.lookup({**scoped, "scope": "team-b"}) is None
client.close()
def test_selected_redis_semantic_runtime_declines_configuration_drift(
redis_stack: tuple[str, str],
semantic_embedding: DeterministicEmbedding,
monkeypatch: pytest.MonkeyPatch,
) -> None:
url, index = redis_stack
facade: Final = semantic_facade(url, index)
resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade))
assert resolver.resolve().kind == "native"
with rebound(facade.cache, "similarity_threshold", 0.5):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
with rebound(facade, "semantic_cache_scope", "end_user"):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
with rebound(facade.cache, "embedding_model", "other-model"):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
with rebound(facade.cache, "_index_name", "other-index"):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
with rebound(facade.cache, "CACHE_KEY_FIELD_NAME", "other-field"):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
def patched_embedding(self: object, prompt: str, metadata: object = None) -> list[float]:
return _semantic_embedding(prompt)
monkeypatch.setattr(RedisSemanticCache, "_get_embedding", patched_embedding)
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
async def test_redis_semantic_explicit_selection_activates_natively(
redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding, monkeypatch: pytest.MonkeyPatch
) -> None:
del semantic_embedding
url, index = redis_stack
facade: Final = activate_native(
Cache(
type=LiteLLMCacheType.REDIS_SEMANTIC,
redis_url=url,
similarity_threshold=0.8,
redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL,
redis_semantic_cache_index_name=index,
)
)
assert_native_runtime(facade)
kwargs: Final = {"model": "gpt-4o", "messages": semantic_messages("name a primary color")}
await facade.async_add_cache({"answer": "blue"}, **kwargs)
assert await facade.async_get_cache(**kwargs) == {"answer": "blue"}

View file

@ -1,253 +0,0 @@
import asyncio
from collections.abc import Callable
from pathlib import Path
from typing import Final, TypeAlias, cast
from urllib.parse import urlparse
from uuid import uuid4
import pytest
from litellm.caching.caching import Cache
from litellm.rust_bridge import _native
from litellm.rust_bridge.response_cache import NativeResponseCacheRuntime, ResponseCacheRuntime
from litellm.types.caching import LiteLLMCacheType
from litellm.types.utils import EmbeddingResponse
from tests.test_litellm_rust.support.cache import activate_native, assert_native_runtime, completion_kwargs
from tests.test_litellm_rust.support.s3_stub import S3Stub
pytestmark: Final = pytest.mark.requires_rust_extension
CacheFactory: TypeAlias = Callable[[], Cache]
@pytest.fixture
def cache_factory(request: pytest.FixtureRequest, tmp_path: Path) -> CacheFactory:
backend: Final = cast(LiteLLMCacheType, request.param)
match backend:
case LiteLLMCacheType.LOCAL:
return lambda: Cache(type=backend)
case LiteLLMCacheType.DISK:
return lambda: Cache(type=backend, disk_cache_dir=str(tmp_path))
case LiteLLMCacheType.REDIS:
parsed: Final = urlparse(cast(str, request.getfixturevalue("redis_url")))
return lambda: Cache(type=backend, host=parsed.hostname, port=str(parsed.port))
case LiteLLMCacheType.S3:
stub: Final = cast(S3Stub, request.getfixturevalue("s3_stub"))
return lambda: Cache(
type=backend,
s3_bucket_name="cache-bucket",
s3_region_name="us-east-1",
s3_endpoint_url=stub.url,
s3_aws_access_key_id="key",
s3_aws_secret_access_key="secret",
s3_path="team",
)
case LiteLLMCacheType.GCS:
return lambda: Cache(type=backend, gcs_bucket_name="bucket", gcs_path="cache/")
case LiteLLMCacheType.REDIS_SEMANTIC:
return lambda: Cache(
type=backend,
redis_url="redis://127.0.0.1:6379",
similarity_threshold=0.8,
redis_semantic_cache_embedding_model="text-embedding-3-small",
)
case LiteLLMCacheType.VALKEY_SEMANTIC:
return lambda: Cache(type=backend, redis_url="redis://127.0.0.1:6390/0", similarity_threshold=0.8)
case _:
raise AssertionError(f"no local factory for {backend}")
ROUND_TRIP_BACKENDS: Final = (
LiteLLMCacheType.LOCAL,
LiteLLMCacheType.DISK,
LiteLLMCacheType.REDIS,
LiteLLMCacheType.S3,
)
SHARED_STORE_BACKENDS: Final = (LiteLLMCacheType.DISK, LiteLLMCacheType.REDIS, LiteLLMCacheType.S3)
@pytest.mark.parametrize(
"cache_factory",
[
LiteLLMCacheType.LOCAL,
LiteLLMCacheType.DISK,
LiteLLMCacheType.REDIS,
LiteLLMCacheType.S3,
LiteLLMCacheType.GCS,
LiteLLMCacheType.REDIS_SEMANTIC,
LiteLLMCacheType.VALKEY_SEMANTIC,
],
indirect=True,
)
def test_legacy_constructor_keeps_python_backends(cache_factory: CacheFactory) -> None:
assert cache_factory()._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor
@pytest.mark.parametrize(
"cache_factory",
[
LiteLLMCacheType.LOCAL,
LiteLLMCacheType.DISK,
LiteLLMCacheType.REDIS,
LiteLLMCacheType.S3,
LiteLLMCacheType.GCS,
LiteLLMCacheType.REDIS_SEMANTIC,
LiteLLMCacheType.VALKEY_SEMANTIC,
],
indirect=True,
)
def test_explicit_selection_activates_the_native_backend(
cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest
) -> None:
assert_native_runtime(activate_native(cache_factory()))
@pytest.mark.parametrize("cache_factory", ROUND_TRIP_BACKENDS, indirect=True)
async def test_facade_storage_calls_round_trip_through_the_native_backend(
cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest
) -> None:
facade: Final = activate_native(cache_factory())
assert_native_runtime(facade)
sync_kwargs: Final = completion_kwargs("sync")
facade.add_cache({"answer": 1}, **sync_kwargs)
assert facade.get_cache(**sync_kwargs) == {"answer": 1}
async_kwargs: Final = completion_kwargs("async")
await facade.async_add_cache({"answer": 2}, **async_kwargs)
assert await facade.async_get_cache(**async_kwargs) == {"answer": 2}
assert facade.get_cache(**completion_kwargs("absent")) is None
async def test_memory_facade_writes_bypass_the_python_backend(monkeypatch: pytest.MonkeyPatch) -> None:
facade: Final = activate_native(Cache(type=LiteLLMCacheType.LOCAL))
assert_native_runtime(facade)
kwargs: Final = completion_kwargs("memory")
facade.add_cache({"answer": 1}, **kwargs)
assert facade.cache.get_cache(facade.get_cache_key(**kwargs)) is None
assert facade.get_cache(**kwargs) == {"answer": 1}
@pytest.mark.parametrize("cache_factory", SHARED_STORE_BACKENDS, indirect=True)
async def test_native_and_python_facades_share_one_wire_format(
cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest
) -> None:
python_facade: Final = cache_factory()
assert python_facade._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor
native_facade: Final = activate_native(cache_factory())
assert_native_runtime(native_facade)
native_written: Final = completion_kwargs("native")
native_facade.add_cache({"writer": "native"}, **native_written)
assert python_facade.get_cache(**native_written) == {"writer": "native"}
python_written: Final = completion_kwargs("python")
python_facade.add_cache({"writer": "python"}, **python_written)
assert native_facade.get_cache(**python_written) == {"writer": "python"}
async_native: Final = completion_kwargs("async-native")
await native_facade.async_add_cache({"writer": "async-native"}, **async_native)
assert await python_facade.async_get_cache(**async_native) == {"writer": "async-native"}
async_python: Final = completion_kwargs("async-python")
await python_facade.async_add_cache({"writer": "async-python"}, **async_python)
assert await native_facade.async_get_cache(**async_python) == {"writer": "async-python"}
@pytest.mark.parametrize("cache_factory", ROUND_TRIP_BACKENDS, indirect=True)
async def test_embedding_pipeline_stores_one_native_entry_per_input(
cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest
) -> None:
facade: Final = activate_native(cache_factory())
assert_native_runtime(facade)
inputs: Final = [f"alpha {uuid4().hex}", f"beta {uuid4().hex}"]
result: Final = EmbeddingResponse(
model="text-embedding-3-small",
data=[
{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]},
{"object": "embedding", "index": 1, "embedding": [0.3, 0.4]},
],
)
await facade.async_add_cache_pipeline(result, model="text-embedding-3-small", input=inputs)
keys: Final = [facade.get_cache_key(model="text-embedding-3-small", input=text) for text in inputs]
assert len(set(keys)) == len(inputs)
for text, expected in zip(inputs, ([0.1, 0.2], [0.3, 0.4]), strict=True):
cached = await facade.async_get_cache(model="text-embedding-3-small", input=text)
assert isinstance(cached, dict)
assert cached["embedding"] == expected
assert await facade.async_get_cache(model="text-embedding-3-small", input=inputs) is None
@pytest.mark.parametrize(
("backend", "settings", "message"),
[
pytest.param(
LiteLLMCacheType.VALKEY_SEMANTIC,
{"redis_url": "rediss://127.0.0.1:6390/0", "similarity_threshold": 0.8},
"native Valkey semantic cache does not support TLS connections",
id="valkey-tls",
),
pytest.param(
LiteLLMCacheType.VALKEY_SEMANTIC,
{"redis_url": "redis://127.0.0.1:6390/0?socket_timeout=1", "similarity_threshold": 0.8},
"native Redis uses fixed socket timeouts; socket_timeout and socket_connect_timeout require Python",
id="valkey-socket-timeout",
),
pytest.param(
LiteLLMCacheType.REDIS_SEMANTIC,
{"redis_url": "rediss://127.0.0.1:6380", "similarity_threshold": 0.8},
"native Redis semantic cache does not support TLS or query options in redis_url",
id="redis-semantic-tls",
),
pytest.param(
LiteLLMCacheType.REDIS_SEMANTIC,
{"redis_url": "redis://127.0.0.1:6379?socket_timeout=1", "similarity_threshold": 0.8},
"native Redis semantic cache does not support TLS or query options in redis_url",
id="redis-semantic-query",
),
],
)
def test_semantic_settings_the_native_client_cannot_honor_decline(
monkeypatch: pytest.MonkeyPatch, backend: LiteLLMCacheType, settings: dict[str, object], message: str
) -> None:
with pytest.raises(_native.RustBridgeDeclined, match=message):
activate_native(Cache(type=backend, **settings))
class _SemanticHit:
"""A native semantic runtime that answers every lookup with one cached response."""
kind: Final = "native"
def lookup_semantic(self, request: object) -> tuple[object, float | None]:
return {"answer": 42}, 0.97
async def async_lookup_semantic(self, request: object) -> tuple[object, float | None]:
return {"answer": 42}, 0.97
@pytest.mark.parametrize("semantic_type", [LiteLLMCacheType.QDRANT_SEMANTIC, LiteLLMCacheType.REDIS_SEMANTIC])
@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"])
def test_native_semantic_hit_stamps_similarity_on_request_metadata(
semantic_type: LiteLLMCacheType, use_async: bool
) -> None:
"""Python semantic backends write `metadata["semantic-similarity"]` on every lookup, and the
facade copies it to the caller's metadata; the native path must report it the same way."""
facade: Final = Cache()
facade.type = semantic_type
facade._native_cache = ResponseCacheRuntime(cast(NativeResponseCacheRuntime, _SemanticHit())) # pyright: ignore[reportPrivateUsage] # the native path under test has no public setter
metadata: Final[dict[str, object]] = {}
kwargs: Final = {
"cache_key": "semantic-key",
"messages": [{"role": "user", "content": "hello"}],
"metadata": metadata,
}
result: Final = asyncio.run(facade.async_get_cache(**kwargs)) if use_async else facade.get_cache(**kwargs)
assert result == {"answer": 42}
assert metadata["semantic-similarity"] == 0.97

View file

@ -1,167 +0,0 @@
import json
import time
from datetime import datetime
from types import SimpleNamespace
from typing import Final, cast
from unittest.mock import Mock
import boto3
import botocore.config
import pytest
from litellm.caching.caching import Cache
from litellm.caching.s3_cache import S3Cache
from litellm.rust_bridge import _native
from litellm.types.caching import LiteLLMCacheType
from tests.test_litellm_rust.support.cache import CacheTestResolver, activate_native, native_runtime, request
from tests.test_litellm_rust.support.isolation import rebound
from tests.test_litellm_rust.support.s3_stub import S3Stub
pytestmark: Final = pytest.mark.requires_rust_extension
def python_s3(url: str) -> S3Cache:
return S3Cache(
s3_bucket_name="cache-bucket",
s3_region_name="us-east-1",
s3_endpoint_url=url,
s3_aws_access_key_id="key",
s3_aws_secret_access_key="secret",
s3_path="team",
)
def s3_facade(url: str) -> Cache:
return Cache(
type=LiteLLMCacheType.S3,
s3_bucket_name="cache-bucket",
s3_region_name="us-east-1",
s3_endpoint_url=url,
s3_aws_access_key_id="key",
s3_aws_secret_access_key="secret",
s3_path="team",
)
async def test_s3_reads_python_entries_and_writes_with_python_metadata(s3_stub: S3Stub) -> None:
python_cache: Final = python_s3(s3_stub.url)
response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}}
python_cache.set_cache("sync:key", {"timestamp": time.time(), "response": response}, ttl=90)
python_cache.set_cache("plain", {"timestamp": time.time(), "response": response})
s3_stub.put_object("team/malformed", b"not a cache entry")
s3_stub.put_object(
"team/expired",
json.dumps({"timestamp": time.time(), "response": response}).encode(),
{"expires": "Thu, 01 Jan 1970 00:00:00 GMT"},
)
binding: Final = native_runtime(s3_facade(s3_stub.url))
assert binding.lookup(request("sync:key")) == response
assert await binding.async_lookup(request("plain")) == response
assert binding.lookup(request("malformed")) is None
assert binding.lookup(request("expired")) is None
assert binding.lookup(request("absent")) is None
binding.store({**request("native:key"), "ttl_seconds": 90.0}, response)
await binding.async_store(request("no_ttl"), response)
stored: Final = s3_stub.objects["team/native/key"]
assert stored.headers["content-type"] == "application/json"
assert stored.headers["content-language"] == "en"
assert stored.headers["content-disposition"] == 'inline; filename="team/native/key.json"'
assert stored.headers["cache-control"] == "immutable, max-age=90, s-maxage=90"
expires: Final = cast(datetime, s3_stub.expires("team/native/key"))
remaining: Final = (expires - datetime.now(expires.tzinfo)).total_seconds()
assert 60 < remaining <= 91
no_ttl: Final = s3_stub.objects["team/no_ttl"]
assert no_ttl.headers["cache-control"] == "immutable, max-age=31536000, s-maxage=31536000"
assert "expires" not in no_ttl.headers
assert python_cache.get_cache("native:key")["response"] == response
partial: Final = await binding.async_lookup_batch([request("native:key"), request("absent"), request("malformed")])
assert partial == {"values": [response, None, None], "missing_indices": [1, 2]}
def test_selected_s3_runtime_declines_backend_mutation(s3_stub: S3Stub) -> None:
facade: Final = Cache(
type=LiteLLMCacheType.S3,
s3_bucket_name="cache-bucket",
s3_region_name="us-east-1",
s3_endpoint_url=s3_stub.url,
s3_aws_access_key_id="key",
s3_aws_secret_access_key="secret",
s3_path="team",
)
activate_native(facade)
resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade))
binding: Final = resolver.resolve()
assert binding.kind == "native"
handler: Final = Mock()
facade.cache.s3_client.meta.events.register("before-call.s3.*", handler)
binding.store(request("native"), {"answer": 1})
assert binding.lookup(request("native")) == {"answer": 1}
assert handler.call_count == 0
assert "team/native" in s3_stub.objects
with rebound(facade.cache, "bucket_name", "other"):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
other_client: Final = boto3.client(
"s3",
region_name="us-east-1",
endpoint_url=s3_stub.url,
aws_access_key_id="key",
aws_secret_access_key="secret",
)
with rebound(facade.cache, "s3_client", other_client):
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
class CustomS3Cache(S3Cache):
pass
subclassed: Final = Cache(
type=LiteLLMCacheType.S3,
s3_bucket_name="cache-bucket",
s3_region_name="us-east-1",
s3_endpoint_url=s3_stub.url,
s3_aws_access_key_id="key",
s3_aws_secret_access_key="secret",
s3_path="team",
)
subclassed.cache = CustomS3Cache(
s3_bucket_name="cache-bucket",
s3_region_name="us-east-1",
s3_endpoint_url=s3_stub.url,
s3_aws_access_key_id="key",
s3_aws_secret_access_key="secret",
s3_path="team",
)
assert CacheTestResolver(SimpleNamespace(cache=subclassed)).resolve().kind == "python_callback"
def test_s3_facade_rejects_configurations_that_require_python(s3_stub: S3Stub) -> None:
unverified: Final = Cache(
type=LiteLLMCacheType.S3,
s3_bucket_name="cache-bucket",
s3_region_name="us-east-1",
s3_endpoint_url="https://s3.example.test",
s3_aws_access_key_id="key",
s3_aws_secret_access_key="secret",
s3_path="team",
s3_verify=False,
)
with pytest.raises(_native.RustBridgeDeclined, match="requires Python"):
native_runtime(unverified)
proxied: Final = Cache(
type=LiteLLMCacheType.S3,
s3_bucket_name="cache-bucket",
s3_region_name="us-east-1",
s3_endpoint_url=s3_stub.url,
s3_aws_access_key_id="key",
s3_aws_secret_access_key="secret",
s3_path="team",
s3_config=botocore.config.Config(proxies={"https": "http://proxy.test"}),
)
with pytest.raises(_native.RustBridgeDeclined, match="requires Python"):
native_runtime(proxied)

View file

@ -1,530 +0,0 @@
import asyncio
import contextvars
import hashlib
import os
import struct
import threading
import time
from collections.abc import Generator, Mapping
from types import SimpleNamespace
from typing import Final, cast
from uuid import uuid4
import pytest
import redis
from litellm.caching.caching import Cache
from litellm.caching.valkey_semantic_cache import ValkeySemanticCache
from litellm.rust_bridge import _native
from litellm.rust_bridge.response_cache import ResponseCacheRuntime
from litellm.types.caching import LiteLLMCacheType
from tests.test_litellm_rust.support.cache import CacheTestResolver, activate_native, native_runtime
pytestmark: Final = pytest.mark.requires_rust_extension
embedding_context: Final = contextvars.ContextVar("embedding_context")
@pytest.fixture
def valkey_url() -> str:
url: Final = os.environ.get("LITELLM_TEST_VALKEY_URL")
if url is None:
pytest.skip("LITELLM_TEST_VALKEY_URL is not set")
return url
@pytest.fixture
def index_name(valkey_url: str) -> Generator[str]:
index: Final = f"litellm_test_{uuid4().hex}"
yield index
client: Final = redis.Redis.from_url(valkey_url)
try:
client.ft(index).dropindex(delete_documents=True)
except redis.ResponseError:
pass
finally:
client.close()
def _request(prompt: str = "semantic cache prompt") -> dict[str, object]:
return {
"key": {"preset": "key"},
"messages": [{"role": "user", "content": prompt}],
}
def _field_request(
prompt: str,
metadata: Mapping[str, object],
*,
namespace: str | None = None,
litellm_metadata: Mapping[str, object] | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> dict[str, object]:
request: Final = {
"key": {
"fields": [
{
"name": "model",
"value": "gpt-4.1",
"api_parameter": True,
"internal_parameter": False,
},
{
"name": "messages",
"value": prompt,
"api_parameter": True,
"internal_parameter": False,
},
],
"namespace": namespace,
},
"messages": [{"role": "user", "content": prompt}],
"metadata": dict(metadata),
}
if litellm_metadata is not None:
request["litellm_metadata"] = dict(litellm_metadata)
if litellm_params is not None:
request["litellm_params"] = dict(litellm_params)
return request
def _facade(
url: str,
index_name: str,
embeddings: Mapping[str, list[float]] | None = None,
*,
namespace: str | None = None,
) -> Cache:
facade: Final = Cache(
type=LiteLLMCacheType.VALKEY_SEMANTIC,
redis_url=url,
similarity_threshold=0.8,
valkey_semantic_cache_index_name=index_name,
namespace=namespace,
)
vectors: Final = embeddings or {"semantic cache prompt": [1.0, 0.0]}
def embed(prompt: str, metadata: Mapping[str, object] | None = None) -> list[float]:
return vectors[prompt]
async def async_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]:
return vectors[prompt]
facade.cache._get_embedding = embed
facade.cache._get_async_embedding = async_embedding
return facade
def test_python_write_native_read(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
backend: Final = cast(ValkeySemanticCache, facade.cache)
response: Final = {"answer": "python"}
backend.set_cache("key", response, messages=_request()["messages"])
binding: Final = native_runtime(facade)
assert binding.lookup(_request()) == response
def test_native_write_python_read(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
backend: Final = cast(ValkeySemanticCache, facade.cache)
binding: Final = native_runtime(facade)
response: Final = {"answer": "native"}
binding.store({**_request(), "ttl_seconds": 2.0}, response)
cached: Final = cast(Mapping[str, object], backend.get_cache("key", messages=_request()["messages"]))
assert cached["response"] == response
async def test_async_lookup_and_store(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
binding: Final = native_runtime(facade)
request: Final = {**_request(), "ttl_seconds": 2.0}
await binding.async_store(request, {"answer": "async"})
assert await binding.async_lookup(request) == {"answer": "async"}
async def test_disabled_cache_controls_skip_async_embedding(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
backend: Final = cast(ValkeySemanticCache, facade.cache)
calls: Final = []
async def fail_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]:
calls.append(prompt)
raise AssertionError("embedding must not run")
backend._get_async_embedding = fail_embedding
binding: Final = native_runtime(facade)
controls: Final = {
"supported_call_type": True,
"configured": True,
"native_backend": True,
"default_on": True,
"caching": True,
"no_cache": False,
"no_store": False,
"use_cache": True,
}
no_read_request: Final = {**_request(), "controls": {**controls, "no_cache": True}}
assert await binding.async_lookup(no_read_request) is None
no_write_request: Final = {**_request(), "controls": {**controls, "no_store": True}}
await binding.async_store(no_write_request, {"answer": "blocked"})
assert calls == []
client: Final = redis.Redis.from_url(valkey_url)
assert list(client.scan_iter(f"{index_name}:*")) == []
client.close()
async def test_async_embedding_runs_inline_in_caller_task(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
backend: Final = cast(ValkeySemanticCache, facade.cache)
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
binding: Final = native_runtime(facade)
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_selected_valkey_runtime_declines_threshold_mutation(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = Cache(
type=LiteLLMCacheType.VALKEY_SEMANTIC,
redis_url=valkey_url,
similarity_threshold=0.8,
valkey_semantic_cache_index_name=index_name,
)
activate_native(facade)
resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade))
assert resolver.resolve().kind == "native"
facade.cache.similarity_threshold = 0.7
with pytest.raises(_native.RustBridgeDeclined):
resolver.resolve()
def test_batch_lookup_is_unsupported(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
binding: Final = native_runtime(facade)
with pytest.raises(NotImplementedError):
binding.lookup_batch([_request()])
def test_ttl_expiry(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
binding: Final = native_runtime(facade)
binding.store({**_request(), "ttl_seconds": 1.0}, {"answer": "expires"})
client: Final = redis.Redis.from_url(valkey_url)
documents: Final = list(client.scan_iter(f"{index_name}:*"))
assert len(documents) == 1
assert client.ttl(documents[0]) > 0
time.sleep(1.5)
assert binding.lookup(_request()) is None
def test_no_ttl_is_persistent_and_python_reads_native_value(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
backend: Final = cast(ValkeySemanticCache, facade.cache)
binding: Final = native_runtime(facade)
response: Final = {"answer": "persistent"}
binding.store(_request(), response)
client: Final = redis.Redis.from_url(valkey_url)
documents: Final = list(client.scan_iter(f"{index_name}:*"))
assert len(documents) == 1
assert client.ttl(documents[0]) == -1
cached: Final = cast(Mapping[str, object], backend.get_cache("key", messages=_request()["messages"]))
assert cached["response"] == response
def test_below_threshold_misses_on_native_and_python(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(
valkey_url,
index_name,
{"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]},
)
backend: Final = cast(ValkeySemanticCache, facade.cache)
binding: Final = native_runtime(facade)
binding.store(_request("prompt A"), {"answer": "A"})
assert binding.lookup(_request("prompt B")) is None
assert backend.get_cache("key", messages=_request("prompt B")["messages"]) is None
def test_malformed_entry_is_a_miss_on_native_and_python(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
backend: Final = cast(ValkeySemanticCache, facade.cache)
client: Final = redis.Redis.from_url(valkey_url)
scope: Final = hashlib.sha256(b"key").hexdigest()
document: Final = f"{index_name}:{scope}:{uuid4().hex}"
client.hset(
document,
mapping={
"litellm_cache_key": scope,
"prompt": "semantic cache prompt",
"response": "not json",
"embedding": struct.pack("<2f", 1.0, 0.0),
},
)
binding: Final = native_runtime(facade)
assert binding.lookup(_request()) is None
assert backend.get_cache("key", messages=_request()["messages"]) is None
def test_mixed_content_parts_match_python_semantic_behavior(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
backend: Final = cast(ValkeySemanticCache, facade.cache)
messages: Final = [{"role": "user", "content": ["raw", {"text": "hello"}]}]
backend.set_cache("key", {"answer": "mixed"}, messages=messages)
assert backend.get_cache("key", messages=messages) is None
binding: Final = native_runtime(facade)
request: Final = {**_request(), "messages": messages}
binding.store(request, {"answer": "mixed"})
assert binding.lookup(request) is None
client: Final = redis.Redis.from_url(valkey_url)
assert list(client.scan_iter(f"{index_name}:*")) == []
client.close()
async def test_async_store_batch_and_lookup(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(
valkey_url,
index_name,
{"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]},
)
backend: Final = cast(ValkeySemanticCache, facade.cache)
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
binding: Final = native_runtime(facade)
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]
def test_subclass_backend_falls_back_to_python(
valkey_url: str,
index_name: str,
) -> None:
class Custom(ValkeySemanticCache):
pass
facade: Final = Cache(
type=LiteLLMCacheType.VALKEY_SEMANTIC,
redis_url=valkey_url,
similarity_threshold=0.8,
valkey_semantic_cache_index_name=index_name,
)
facade.cache = Custom(redis_url=valkey_url, similarity_threshold=0.8, index_name=index_name)
resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade))
assert resolver.resolve().kind == "python_callback"
def test_field_key_matches_python_semantic_scope(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]})
metadata: Final = {"user_api_key": "k1"}
expected: Final = facade.get_cache_key(
model="gpt-4.1",
messages=[{"role": "user", "content": "semantic cache prompt"}],
metadata=metadata,
)
binding: Final = native_runtime(facade)
binding.store(_field_request("semantic cache prompt", metadata), {"answer": "scoped"})
client: Final = redis.Redis.from_url(valkey_url)
documents: Final = list(client.scan_iter(f"{index_name}:*"))
assert len(documents) == 1
document_parts: Final = documents[0].decode().split(":")
assert document_parts[1] == hashlib.sha256(expected.encode()).hexdigest()
client.close()
def test_field_key_reads_all_python_tenant_metadata_sources(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]})
params_metadata: Final = {"user_api_key_team_id": "team-from-params"}
expected: Final = facade.get_cache_key(
model="gpt-4.1",
messages=[{"role": "user", "content": "semantic cache prompt"}],
metadata={},
litellm_params={"metadata": params_metadata},
)
binding: Final = native_runtime(facade)
binding.store(
_field_request(
"semantic cache prompt",
{},
litellm_params={"metadata": params_metadata},
),
{"answer": "params"},
)
client: Final = redis.Redis.from_url(valkey_url)
documents: Final = list(client.scan_iter(f"{index_name}:*"))
assert len(documents) == 1
document_parts: Final = documents[0].decode().split(":")
assert document_parts[1] == hashlib.sha256(expected.encode()).hexdigest()
client.close()
assert (
binding.lookup(
_field_request(
"semantic cache prompt",
{},
litellm_metadata={"user_api_key_team_id": "team-from-litellm"},
)
)
is None
)
def test_namespace_isolates_semantic_entries(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(
valkey_url,
index_name,
{"semantic cache prompt": [1.0, 0.0]},
namespace="team-a",
)
binding: Final = native_runtime(facade)
team_a: Final = _field_request("semantic cache prompt", {}, namespace="team-a")
team_b: Final = _field_request("semantic cache prompt", {}, namespace="team-b")
binding.store(team_a, {"answer": "team-a"})
assert binding.lookup(team_b) is None
assert binding.lookup(team_a) == {"answer": "team-a"}
cached: Final = cast(
Mapping[str, object],
facade.get_cache(
model="gpt-4.1",
messages=[{"role": "user", "content": "semantic cache prompt"}],
),
)
assert cached == {"answer": "team-a"}
def test_field_key_isolates_tenant_scope(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]})
binding: Final = native_runtime(facade)
binding.store(
_field_request("semantic cache prompt", {"user_api_key": "k1"}),
{"answer": "tenant one"},
)
assert binding.lookup(_field_request("semantic cache prompt", {"user_api_key": "k2"})) is None
assert binding.lookup(_field_request("semantic cache prompt", {"user_api_key": "k1"})) == {"answer": "tenant one"}
def test_tls_valkey_facade_falls_back_to_python(
index_name: str,
) -> None:
facade: Final = Cache(
type=LiteLLMCacheType.VALKEY_SEMANTIC,
redis_url="rediss://127.0.0.1:6390/0",
similarity_threshold=0.8,
valkey_semantic_cache_index_name=index_name,
)
resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade))
assert resolver.resolve().kind == "python_callback"
async def test_ping_maps_unsupported_native_operation_to_not_implemented(
valkey_url: str,
index_name: str,
) -> None:
facade: Final = _facade(valkey_url, index_name)
binding: Final = native_runtime(facade)
with pytest.raises(NotImplementedError):
await binding.ping()
async def test_explicit_selection_activates_the_facade_natively(
valkey_url: str,
index_name: str,
monkeypatch: pytest.MonkeyPatch,
) -> None:
facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]})
facade._native_cache = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade)) # pyright: ignore[reportPrivateUsage] # explicitly select the runtime under test
runtime: Final = facade._native_cache # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor
assert isinstance(runtime, ResponseCacheRuntime)
assert runtime.kind == "native"
kwargs: Final = {"model": "gpt-4o", "messages": _request()["messages"]}
await facade.async_add_cache({"answer": "valkey"}, **kwargs)
assert await facade.async_get_cache(**kwargs) == {"answer": "valkey"}

View file

@ -1,73 +1,25 @@
from __future__ import annotations
from collections.abc import AsyncIterator, Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Literal, Protocol, TypeAlias
from uuid import uuid4
from typing import Final, Literal
from pydantic import BaseModel, TypeAdapter
import litellm
from litellm.caching.caching import Cache
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
from litellm.rust_bridge import _native, runtime
from litellm.rust_bridge import runtime
from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule
from litellm.rust_bridge.chat_completions.entrypoints import NATIVE_ACOMPLETION, LiteLLMChatCompletionsRequest
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.dispatch import call_hook
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, LiteLLMMessagesRequest
from litellm.rust_bridge.response_cache import ResponseCacheRuntime
from litellm.rust_bridge.responses.entrypoints import NATIVE_ARESPONSES, LiteLLMResponsesRequest
from litellm.types.utils import ModelResponse
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_EVENTS, MESSAGES_MODEL, MESSAGES_RESPONSE
from tests.test_litellm_rust.test_inference import RESPONSES_MODEL, RESPONSES_RESPONSE
CacheRuntime: TypeAlias = _native._ResponseCacheRuntime # pyright: ignore[reportPrivateUsage] # private runtime under test
class CacheNamespace(Protocol):
@property
def cache(self) -> object: ...
@dataclass(frozen=True, slots=True)
class CacheTestResolver:
namespace: CacheNamespace
def resolve(self) -> CacheRuntime:
return CacheRuntime.from_selected(self.namespace.cache)
class CacheLookup(Protocol):
def get_cache(self, **kwargs: object) -> object: ...
def flush_cache(self) -> object: ...
def request(key: str = "key") -> dict[str, object]:
return {"key": {"preset": key}}
def native_runtime(facade: Cache) -> CacheRuntime:
return CacheRuntime.from_cache(facade)
def activate_native(facade: Cache) -> Cache:
facade._native_cache = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade)) # pyright: ignore[reportPrivateUsage] # explicitly select the runtime under test
return facade
def assert_native_runtime(facade: Cache) -> ResponseCacheRuntime:
runtime: Final = facade._native_cache # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor
assert isinstance(runtime, ResponseCacheRuntime)
assert runtime.kind == "native"
return runtime
def completion_kwargs(label: str) -> dict[str, object]:
return {"model": "gpt-4o", "messages": [{"role": "user", "content": f"{label} {uuid4().hex}"}]}
def payload(value: object) -> object:
if isinstance(value, ModelResponse):

View file

@ -1,112 +0,0 @@
"""In-process path-style S3 stub for native cache parity tests."""
import threading
from dataclasses import dataclass, field
from email.utils import parsedate_to_datetime
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Final
from urllib.parse import unquote, urlsplit
_STORED_HEADERS: Final = (
"cache-control",
"content-type",
"content-language",
"content-disposition",
"expires",
)
@dataclass
class S3Object:
body: bytes
headers: dict[str, str] = field(default_factory=dict)
class S3Stub:
"""Minimal path-style S3 endpoint serving PUT and GET object operations."""
def __init__(self) -> None:
self._objects: dict[str, S3Object] = {}
stub: Final = self
class Handler(BaseHTTPRequestHandler):
def _key(self) -> str:
parts: Final = urlsplit(self.path).path.lstrip("/").split("/", 1)
return unquote(parts[1]) if len(parts) == 2 else ""
def _read_body(self) -> bytes:
transfer: Final = self.headers.get("transfer-encoding", "")
if "chunked" not in transfer:
return self.rfile.read(int(self.headers.get("content-length", 0)))
chunks: Final = bytearray()
while True:
size = int(self.rfile.readline().split(b";")[0].strip(), 16)
if size == 0:
while self.rfile.readline().strip():
pass
return bytes(chunks)
chunks.extend(self.rfile.read(size))
self.rfile.readline()
def do_PUT(self) -> None:
body: Final = self._read_body()
headers: Final = {name: self.headers[name] for name in _STORED_HEADERS if name in self.headers}
stub._objects = {**stub._objects, self._key(): S3Object(body=body, headers=headers)}
self.send_response(200)
self.send_header("ETag", '"stub"')
self.send_header("Content-Length", "0")
self.end_headers()
def do_HEAD(self) -> None:
self._object(send_body=False)
def do_GET(self) -> None:
self._object(send_body=True)
def _object(self, send_body: bool) -> None:
entry: Final = stub._objects.get(self._key())
if entry is None:
self.send_response(404)
self.send_header("Content-Type", "application/xml")
body: Final = b'<?xml version="1.0" encoding="UTF-8"?><Error><Code>NoSuchKey</Code></Error>'
self.send_header("Content-Length", str(len(body)))
self.end_headers()
if send_body:
self.wfile.write(body)
return
self.send_response(200)
for name, value in entry.headers.items():
self.send_header(name, value)
self.send_header("ETag", '"stub"')
self.send_header("Content-Length", str(len(entry.body)))
self.end_headers()
if send_body:
self.wfile.write(entry.body)
def log_message(self, format: str, *args: object) -> None:
pass
self._server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
self._worker: Final = threading.Thread(target=self._server.serve_forever, daemon=True)
self._worker.start()
@property
def url(self) -> str:
host, port = self._server.server_address[:2]
return f"http://{host}:{port}"
@property
def objects(self) -> dict[str, S3Object]:
return self._objects
def put_object(self, key: str, body: bytes, headers: dict[str, str] | None = None) -> None:
self._objects = {**self._objects, key: S3Object(body=body, headers=headers or {})}
def expires(self, key: str) -> object:
header: Final = self._objects[key].headers.get("expires")
return parsedate_to_datetime(header) if header else None
def close(self) -> None:
self._server.shutdown()
self._server.server_close()
self._worker.join(timeout=5)