From 5f9eff55f8928e520dec5651695d7bfc98b0090c Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 20:58:19 -0700 Subject: [PATCH] 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 Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/Cargo.lock | 10 - litellm-rust/crates/python-bridge/Cargo.toml | 10 - .../crates/python-bridge/src/cache/AGENTS.md | 12 +- .../crates/python-bridge/src/cache/future.rs | 13 - .../crates/python-bridge/src/cache/mod.rs | 3 - .../python-bridge/src/cache/native/AGENTS.md | 8 +- .../src/cache/native/activation.rs | 112 -- .../python-bridge/src/cache/native/backend.rs | 690 ------- .../python-bridge/src/cache/native/config.rs | 1715 ----------------- .../src/cache/native/embedder.rs | 129 -- .../python-bridge/src/cache/native/facade.rs | 502 ----- .../src/cache/native/identity.rs | 510 ----- .../python-bridge/src/cache/native/mod.rs | 8 - .../python-bridge/src/cache/native/request.rs | 258 --- .../src/cache/native/semantic.rs | 211 -- .../python-bridge/src/cache/native/v2.rs | 34 +- .../python-bridge/src/cache/python/AGENTS.md | 2 +- .../src/cache/python/callback.rs | 165 -- .../python-bridge/src/cache/python/mod.rs | 2 - .../crates/python-bridge/src/cache/runtime.rs | 385 ---- litellm-rust/crates/python-bridge/src/lib.rs | 2 - litellm/caching/caching.py | 84 +- litellm/rust_bridge/_native.pyi | 59 - litellm/rust_bridge/response_cache.py | 146 -- tests/test_litellm_rust/cache/conftest.py | 11 - .../cache/test_azure_blob.py | 165 -- tests/test_litellm_rust/cache/test_disk.py | 115 -- tests/test_litellm_rust/cache/test_facade.py | 356 ---- tests/test_litellm_rust/cache/test_gcs.py | 46 - .../cache/test_qdrant_semantic.py | 254 --- tests/test_litellm_rust/cache/test_redis.py | 212 -- .../cache/test_redis_semantic.py | 560 ------ tests/test_litellm_rust/cache/test_rollout.py | 253 --- tests/test_litellm_rust/cache/test_s3.py | 167 -- .../cache/test_valkey_semantic.py | 530 ----- tests/test_litellm_rust/support/cache.py | 52 +- tests/test_litellm_rust/support/s3_stub.py | 112 -- 37 files changed, 28 insertions(+), 7875 deletions(-) delete mode 100644 litellm-rust/crates/python-bridge/src/cache/future.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/native/activation.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/native/backend.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/native/config.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/native/embedder.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/native/facade.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/native/identity.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/native/request.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/native/semantic.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/python/callback.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/runtime.rs delete mode 100644 litellm/rust_bridge/response_cache.py delete mode 100644 tests/test_litellm_rust/cache/test_azure_blob.py delete mode 100644 tests/test_litellm_rust/cache/test_disk.py delete mode 100644 tests/test_litellm_rust/cache/test_facade.py delete mode 100644 tests/test_litellm_rust/cache/test_gcs.py delete mode 100644 tests/test_litellm_rust/cache/test_qdrant_semantic.py delete mode 100644 tests/test_litellm_rust/cache/test_redis.py delete mode 100644 tests/test_litellm_rust/cache/test_redis_semantic.py delete mode 100644 tests/test_litellm_rust/cache/test_rollout.py delete mode 100644 tests/test_litellm_rust/cache/test_s3.py delete mode 100644 tests/test_litellm_rust/cache/test_valkey_semantic.py delete mode 100644 tests/test_litellm_rust/support/s3_stub.py diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 2461cb17978..83d3a289c42 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -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", ] diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 7458e64d374..105241e1440 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -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" diff --git a/litellm-rust/crates/python-bridge/src/cache/AGENTS.md b/litellm-rust/crates/python-bridge/src/cache/AGENTS.md index 63223a6816c..7a692670eba 100644 --- a/litellm-rust/crates/python-bridge/src/cache/AGENTS.md +++ b/litellm-rust/crates/python-bridge/src/cache/AGENTS.md @@ -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 diff --git a/litellm-rust/crates/python-bridge/src/cache/future.rs b/litellm-rust/crates/python-bridge/src/cache/future.rs deleted file mode 100644 index 242a23242cc..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/future.rs +++ /dev/null @@ -1,13 +0,0 @@ -use litellm_host_python::{ready_future, to_py}; -use pyo3::prelude::*; - -pub(super) fn ready_none(py: Python<'_>) -> PyResult> { - ready_value(py, &()) -} - -pub(super) fn ready_value<'py, T: serde::Serialize>( - py: Python<'py>, - value: &T, -) -> PyResult> { - ready_future(py, to_py(py, value)?.bind(py)) -} diff --git a/litellm-rust/crates/python-bridge/src/cache/mod.rs b/litellm-rust/crates/python-bridge/src/cache/mod.rs index 179f16c4a1f..5a8f9bcf94f 100644 --- a/litellm-rust/crates/python-bridge/src/cache/mod.rs +++ b/litellm-rust/crates/python-bridge/src/cache/mod.rs @@ -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; diff --git a/litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md b/litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md index 65804a18564..09a34d188cd 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md +++ b/litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md @@ -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 diff --git a/litellm-rust/crates/python-bridge/src/cache/native/activation.rs b/litellm-rust/crates/python-bridge/src/cache/native/activation.rs deleted file mode 100644 index 3ac038d2c39..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/native/activation.rs +++ /dev/null @@ -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 { - 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)) -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native/backend.rs b/litellm-rust/crates/python-bridge/src/cache/native/backend.rs deleted file mode 100644 index 1151ed5cc9d..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/native/backend.rs +++ /dev/null @@ -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, -} - -/// An exact-match backend behind one pointer, with the identity its facade must reproduce. -pub(in crate::cache) struct ExactService { - cache: Arc, - probe: Option>, - buffer: Option, - identity: BackendIdentity, -} - -#[derive(Clone)] -pub(in crate::cache) enum NativeResponseCache { - Exact(Arc), - ValkeySemantic { - cache: Arc>>, - embedder: PythonEmbedder, - scope: String, - }, - RedisSemantic { - cache: Arc>>, - embedder: PythonEmbedder, - }, - QdrantSemantic(Arc>>), -} - -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, - namespace: Option, - ) -> Result { - 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) -> Result { - 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) -> 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 { - 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(cache: ResponseCache, identity: BackendIdentity) -> Self - where - ResponseCache: ExactResponseCache + 'static, - B: litellm_cache::BaseCache, - 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(cache: ResponseCache, identity: BackendIdentity) -> Self - where - ResponseCache: ExactResponseCache + ConnectionProbe + 'static, - B: litellm_cache::BaseCache, - B::Context: Default + PartialEq, - { - let cache = Arc::new(cache); - Self::exact_service(cache.clone(), Some(cache), identity) - } - - fn exact_service( - cache: Arc, - probe: Option>, - 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 { - 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 { - 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 { - 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) -> 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> { - 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 { - 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> { - 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, 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, 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 { - 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, 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, 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> { - 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> { - 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> { - 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 { - 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> { - 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 { - 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 { - 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, - pub(in crate::cache) Option, -); - -impl From> for SemanticReply { - fn from(lookup: SemanticLookup) -> Self { - Self(lookup.value, lookup.similarity) - } -} - -fn exact_lookup(value: Option) -> SemanticLookup { - SemanticLookup { - value, - similarity: None, - } -} - -/// Python's Redis and Valkey semantic caches catch every lookup failure and stamp `0.0`. -fn redis_family( - lookup: Result, Error>, -) -> Result, Error> { - Ok(lookup.unwrap_or_else(|_| SemanticLookup::miss(Some(0.0)))) -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native/config.rs b/litellm-rust/crates/python-bridge/src/cache/native/config.rs deleted file mode 100644 index aa1cb1c70fc..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/native/config.rs +++ /dev/null @@ -1,1715 +0,0 @@ -use std::{path::PathBuf, time::Duration}; - -use litellm_auth_aws::AwsAuthConfig; -use litellm_cache::CacheType; -use litellm_cache_qdrant_semantic::{OpenAiEmbedderConfig, QdrantSemanticConfig, Quantization}; -use litellm_cache_redis::{RedisNode, RedisTopology}; -use litellm_cache_s3::{S3CacheConfig, S3Endpoint}; -use pyo3::{ - exceptions::{PyAttributeError, PyTypeError, PyValueError}, - prelude::*, - types::{PyAny, PyBool, PyDict, PyList, PyString}, -}; - -use super::{backend::NativeResponseCache, identity::BackendIdentity, request::duration}; - -pub(super) struct CachePolicy { - pub(super) redis_flush_size: Option, - pub(super) semantic_cache_scope: String, -} - -pub(super) struct MemoryCacheConfig { - pub(super) default_ttl: Duration, - pub(super) capacity: usize, - pub(super) max_entry_bytes: usize, -} - -pub(super) struct DiskCacheConfig { - pub(super) directory: PathBuf, -} - -#[derive(Debug, PartialEq)] -pub(super) enum RedisProtocol { - Resp2, - Resp3, -} - -#[derive(Debug, PartialEq)] -pub(super) enum CertificateRequirement { - None, - Optional, - Required, -} - -pub(super) struct RedisTlsConfig { - pub(super) certificate_requirement: CertificateRequirement, - pub(super) check_hostname: bool, - pub(super) ca_certificate: Option, - pub(super) ca_data: Option, - pub(super) client_certificate: Option, - pub(super) client_key: Option, -} - -pub(super) struct RedisConnectionConfig { - pub(super) host: String, - pub(super) port: u16, - pub(super) database: i64, - pub(super) username: Option, - pub(super) password: Option, - pub(super) protocol: RedisProtocol, - pub(super) pool_size: usize, - pub(super) read_timeout: Option, - pub(super) connect_timeout: Option, - pub(super) socket_keepalive: Option, - pub(super) health_check_interval: Duration, - pub(super) client_name: Option, - pub(super) tls: Option, -} - -pub(super) struct RedisCacheConfig { - pub(super) default_ttl: Duration, - pub(super) namespace: Option, - pub(super) flush_size: usize, - pub(super) topology: RedisTopology, - pub(super) connection: RedisConnectionConfig, -} - -#[derive(Debug, PartialEq)] -pub(super) struct GcsCacheConfig { - pub(super) bucket_name: String, - pub(super) key_prefix: String, - pub(super) path_service_account: Option, -} - -pub(super) struct AzureBlobCacheConfig { - pub(super) account_url: String, - pub(super) container: String, -} - -pub(super) struct RedisSemanticCacheConfig { - pub(super) redis_url: String, - pub(super) index_name: String, - pub(super) similarity_threshold: f64, -} - -struct RedisClientProjection<'py> { - topology: RedisTopology, - host: String, - port: u16, - pool_size: usize, - resolved: Bound<'py, PyDict>, - tls: Option, -} - -const REDIS_PY_DEFAULT_MAX_CONNECTIONS: usize = 1 << 31; - -/// The read and write timeout every native Redis connection uses, which is also `RedisCache`'s -/// default `socket_timeout`. -const NATIVE_REDIS_SOCKET_TIMEOUT: Duration = Duration::from_secs(5); - -pub(super) struct ValkeySemanticCacheConfig { - pub(super) similarity_threshold: f64, - pub(super) index_name: String, - pub(super) connection: RedisConnectionConfig, -} - -pub(super) struct QdrantSemanticCacheConfig { - pub(super) grpc_url: String, - pub(super) api_key: Option, - pub(super) collection_name: String, - pub(super) similarity_threshold: f64, - pub(super) vector_size: u64, - pub(super) embedding: OpenAiEmbedderConfig, - pub(super) quantization: Quantization, -} - -impl QdrantSemanticCacheConfig { - pub(super) fn to_qdrant_config(&self) -> QdrantSemanticConfig { - QdrantSemanticConfig { - collection_name: self.collection_name.clone(), - similarity_threshold: self.similarity_threshold, - vector_size: self.vector_size, - quantization: self.quantization.clone(), - } - } -} - -impl RedisTlsConfig { - /// Whether redis-rs with rustls behaves like this redis-py `SSLConnection`: it verifies the - /// certificate chain against the system roots and always checks the hostname. - fn native(&self) -> Result<(), UnsupportedCacheConfig> { - if self.ca_certificate.is_some() - || self.ca_data.is_some() - || self.client_certificate.is_some() - || self.client_key.is_some() - { - return Err(UnsupportedCacheConfig::RedisTlsCertificates); - } - if self.certificate_requirement == CertificateRequirement::None || !self.check_hostname { - return Err(UnsupportedCacheConfig::RedisTlsVerification); - } - Ok(()) - } -} - -impl RedisConnectionConfig { - /// The redis-rs URL for this connection, or the first setting the native client cannot - /// honor. The native pool and socket timeouts are fixed, so only redis-py's unbounded pool, - /// its unset timeouts and `RedisCache`'s five-second `socket_timeout` map onto them. - pub(super) fn native_url(&self) -> Result { - if self.pool_size != REDIS_PY_DEFAULT_MAX_CONNECTIONS { - return Err(UnsupportedCacheConfig::RedisPoolSize); - } - if self - .read_timeout - .is_some_and(|timeout| timeout != NATIVE_REDIS_SOCKET_TIMEOUT) - || self.connect_timeout.is_some() - { - return Err(UnsupportedCacheConfig::RedisTimeout); - } - if self.socket_keepalive == Some(true) { - return Err(UnsupportedCacheConfig::RedisKeepalive); - } - if !self.health_check_interval.is_zero() { - return Err(UnsupportedCacheConfig::RedisHealthCheck); - } - if self.client_name.is_some() { - return Err(UnsupportedCacheConfig::RedisClientName); - } - let scheme = match &self.tls { - None => "redis", - Some(tls) => { - tls.native()?; - "rediss" - } - }; - let host = if self.host.contains(':') { - format!("[{}]", self.host) - } else { - self.host.clone() - }; - let mut url = url::Url::parse(&format!( - "{scheme}://{host}:{}/{}", - self.port, self.database - )) - .map_err(|_| UnsupportedCacheConfig::RedisConnection)?; - if let Some(username) = &self.username { - url.set_username(username) - .map_err(|()| UnsupportedCacheConfig::RedisConnection)?; - } - if let Some(password) = &self.password { - url.set_password(Some(password)) - .map_err(|()| UnsupportedCacheConfig::RedisConnection)?; - } - if self.protocol == RedisProtocol::Resp3 { - url.set_query(Some("protocol=resp3")); - } - Ok(url.into()) - } -} - -impl RedisSemanticCacheConfig { - /// redisvl hands `redis_url` to redis-py, which reads TLS and socket options from the URL; - /// redis-rs ignores those, so only a plain URL keeps its meaning. - pub(super) fn native_url(&self) -> Result<&str, UnsupportedCacheConfig> { - let url = url::Url::parse(&self.redis_url) - .map_err(|_| UnsupportedCacheConfig::RedisSemanticUrl)?; - if !matches!(url.scheme(), "redis" | "unix") || url.query().is_some() { - return Err(UnsupportedCacheConfig::RedisSemanticUrl); - } - Ok(&self.redis_url) - } -} - -pub(super) enum CacheBackendConfig { - Memory(MemoryCacheConfig), - Redis(Box), - S3(Box), - Gcs(GcsCacheConfig), - ValkeySemantic(Box), - Disk(DiskCacheConfig), - AzureBlob(AzureBlobCacheConfig), - RedisSemantic(Box), - QdrantSemantic(Box), -} - -pub(in crate::cache) struct NativeCacheConfig { - pub(super) policy: CachePolicy, - pub(super) backend: CacheBackendConfig, -} - -pub(in crate::cache) enum UnsupportedCacheConfig { - Backend, - RedisTopology, - RedisCredentials, - RedisConnection, - RedisOption, - S3Client, - S3Credentials, - S3Option, - GcsBucket, - DiskStore, - QdrantEndpoint, - SemanticEmbedding, - RedisPoolSize, - RedisTimeout, - RedisKeepalive, - RedisHealthCheck, - RedisClientName, - RedisTlsCertificates, - RedisTlsVerification, - RedisSemanticUrl, - ValkeyTls, -} - -impl UnsupportedCacheConfig { - pub(in crate::cache) fn message(&self) -> &'static str { - match self { - Self::Backend => "native cache backend is not implemented", - Self::RedisTopology => "native Redis topology is not implemented", - Self::RedisCredentials => "native Redis credentials require Python", - Self::RedisConnection => "native Redis connection type is not implemented", - Self::RedisOption => "native Redis configuration requires Python", - Self::S3Client => "native S3 client type is not implemented", - Self::S3Credentials => "native S3 credentials require Python", - Self::S3Option => "native S3 configuration requires Python", - Self::GcsBucket => "native GCS cache requires a configured bucket name", - Self::DiskStore => "native disk cache requires the built-in diskcache store", - Self::QdrantEndpoint => { - "native Qdrant requires the default REST port so the gRPC port can be derived" - } - Self::SemanticEmbedding => "native semantic embedding requires Python", - Self::RedisPoolSize => { - "native Redis uses a fixed connection pool; max_connections requires Python" - } - Self::RedisTimeout => { - "native Redis uses fixed socket timeouts; socket_timeout and \ - socket_connect_timeout require Python" - } - Self::RedisKeepalive => "native Redis does not support socket_keepalive", - Self::RedisHealthCheck => "native Redis does not support health_check_interval", - Self::RedisClientName => "native Redis does not support client_name", - Self::RedisTlsCertificates => { - "native Redis TLS does not support ssl_ca_certs, ssl_ca_data, ssl_certfile or \ - ssl_keyfile" - } - Self::RedisTlsVerification => { - "native Redis TLS always verifies the certificate and hostname; \ - ssl_cert_reqs=none and ssl_check_hostname=false require Python" - } - Self::RedisSemanticUrl => { - "native Redis semantic cache does not support TLS or query options in redis_url" - } - Self::ValkeyTls => "native Valkey semantic cache does not support TLS connections", - } - } -} - -pub(in crate::cache) enum CacheConfigProjection { - Native(Box), - Unsupported(UnsupportedCacheConfig), -} - -impl NativeCacheConfig { - #[inline(never)] - pub(in crate::cache) fn project(facade: &Bound<'_, PyAny>) -> PyResult { - let backend_name = facade.getattr("type")?.extract::()?; - let policy = CachePolicy { - redis_flush_size: facade - .getattr("redis_flush_size")? - .extract::>()?, - semantic_cache_scope: facade - .getattr("semantic_cache_scope")? - .extract::()?, - }; - let backend = facade.getattr("cache")?; - match CacheType::from_python_name(&backend_name) { - Some(CacheType::Local) => project_memory(&backend).map(|backend| { - CacheConfigProjection::Native(Box::new(Self { - policy, - backend: CacheBackendConfig::Memory(backend), - })) - }), - Some(CacheType::Redis) => match project_redis(&backend)? { - Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { - policy, - backend: CacheBackendConfig::Redis(Box::new(backend)), - }))), - Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), - }, - Some(CacheType::S3) => match project_s3(&backend)? { - Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { - policy, - backend: CacheBackendConfig::S3(Box::new(backend)), - }))), - Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), - }, - Some(CacheType::Gcs) => match project_gcs(&backend)? { - Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { - policy, - backend: CacheBackendConfig::Gcs(backend), - }))), - Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), - }, - Some(CacheType::ValkeySemantic) => match project_valkey_semantic(&backend)? { - Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { - policy, - backend: CacheBackendConfig::ValkeySemantic(Box::new(backend)), - }))), - Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), - }, - Some(CacheType::Disk) => match project_disk(&backend)? { - Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { - policy, - backend: CacheBackendConfig::Disk(backend), - }))), - Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), - }, - Some(CacheType::QdrantSemantic) => match project_qdrant_semantic(&backend)? { - Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { - policy, - backend: CacheBackendConfig::QdrantSemantic(Box::new(backend)), - }))), - Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), - }, - Some(CacheType::AzureBlob) => project_azure_blob(&backend).map(|backend| { - CacheConfigProjection::Native(Box::new(Self { - policy, - backend: CacheBackendConfig::AzureBlob(backend), - })) - }), - Some(CacheType::RedisSemantic) => project_redis_semantic(&backend).map(|backend| { - CacheConfigProjection::Native(Box::new(Self { - policy, - backend: CacheBackendConfig::RedisSemantic(Box::new(backend)), - })) - }), - None => Ok(CacheConfigProjection::Unsupported( - UnsupportedCacheConfig::Backend, - )), - } - } - - pub(super) fn service_mismatch(&self, service: &NativeResponseCache) -> Option<&'static str> { - self.backend.identity().mismatch(&service.identity()) - } -} - -impl CacheBackendConfig { - /// The identity a native backend must have for this facade configuration to describe it. - pub(super) fn identity(&self) -> BackendIdentity { - match self { - Self::Memory(config) => BackendIdentity::Memory { - capacity: config.capacity, - max_entry_bytes: Some(config.max_entry_bytes), - default_ttl: Some(config.default_ttl), - }, - Self::Redis(config) => BackendIdentity::Redis { - topology: config.topology.clone(), - namespace: config.namespace.clone(), - default_ttl: Some(config.default_ttl), - }, - Self::S3(config) => BackendIdentity::S3 { - bucket: config.bucket.clone(), - key_prefix: config.key_prefix.clone(), - region: config.region.clone(), - endpoint: config - .endpoint - .as_ref() - .map(|endpoint| endpoint.url.clone()), - }, - Self::Gcs(config) => BackendIdentity::Gcs { - bucket_name: config.bucket_name.clone(), - key_prefix: config.key_prefix.clone(), - path_service_account: config.path_service_account.clone(), - }, - Self::ValkeySemantic(config) => BackendIdentity::ValkeySemantic { - index_name: config.index_name.clone(), - similarity_threshold: config.similarity_threshold, - }, - Self::Disk(config) => BackendIdentity::Disk { - directory: config.directory.clone(), - }, - Self::AzureBlob(config) => BackendIdentity::AzureBlob { - account_url: config.account_url.clone(), - container: config.container.clone(), - }, - Self::RedisSemantic(config) => BackendIdentity::RedisSemantic { - index_name: config.index_name.clone(), - similarity_threshold: config.similarity_threshold as f32, - }, - Self::QdrantSemantic(config) => BackendIdentity::QdrantSemantic { - collection_name: config.collection_name.clone(), - similarity_threshold: config.similarity_threshold, - vector_size: config.vector_size, - embedding_model: config.embedding.model.clone(), - }, - } - } -} - -#[inline(never)] -fn project_qdrant_semantic( - backend: &Bound<'_, PyAny>, -) -> PyResult> { - let rest_url = backend.getattr("qdrant_api_base")?.extract::()?; - let parsed = match url::Url::parse(&rest_url) { - Ok(value) => value, - Err(_) => return Ok(Err(UnsupportedCacheConfig::QdrantEndpoint)), - }; - if !matches!(parsed.scheme(), "http" | "https") - || (!parsed.path().is_empty() && parsed.path() != "/") - || parsed.query().is_some() - || parsed.host_str().is_none() - || parsed.port() != Some(6333) - { - return Ok(Err(UnsupportedCacheConfig::QdrantEndpoint)); - } - let mut grpc_url = parsed; - if grpc_url.set_port(Some(6334)).is_err() { - return Ok(Err(UnsupportedCacheConfig::QdrantEndpoint)); - } - grpc_url.set_path(""); - grpc_url.set_query(None); - - if optional_attribute(backend, "embedding_max_input_tokens")? - .is_some_and(|value| !value.is_none()) - { - return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); - } - let configured_model = backend.getattr("embedding_model")?.extract::()?; - let embedding_model = configured_model - .strip_prefix("openai/") - .unwrap_or(&configured_model) - .to_owned(); - if !embedding_model.starts_with("text-embedding-") { - return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); - } - let proxy_server = py_sys_module(backend.py())?; - if let Some(proxy_server) = proxy_server { - let router = proxy_server.getattr("llm_router")?; - let model_list = proxy_server.getattr("llm_model_list")?; - let embedding_router = backend.py().import("litellm.caching._embedding_router")?; - if !embedding_router - .getattr("resolve_embedding_router")? - .call1((configured_model.as_str(), router, model_list))? - .is_none() - { - return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); - } - } - let litellm = backend.py().import("litellm")?; - for name in ["api_key", "openai_key", "api_base"] { - if !litellm.getattr(name)?.is_none() { - return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); - } - } - let Ok(embedding_api_key) = std::env::var("OPENAI_API_KEY") else { - return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); - }; - if embedding_api_key.is_empty() { - return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); - } - let embedding_api_base = std::env::var("OPENAI_BASE_URL") - .or_else(|_| std::env::var("OPENAI_API_BASE")) - .unwrap_or_else(|_| "https://api.openai.com/v1".to_owned()); - let timeout = optional_attribute(backend, "embedding_timeout")? - .map(|value| value.extract::>()) - .transpose()? - .flatten() - .map(duration) - .transpose()?; - Ok(Ok(QdrantSemanticCacheConfig { - grpc_url: grpc_url.to_string().trim_end_matches('/').to_owned(), - api_key: optional_string(backend.getattr("qdrant_api_key")?)?, - collection_name: backend.getattr("collection_name")?.extract()?, - similarity_threshold: backend.getattr("similarity_threshold")?.extract()?, - vector_size: backend.getattr("vector_size")?.extract::()?, - embedding: OpenAiEmbedderConfig { - api_base: embedding_api_base, - api_key: embedding_api_key, - model: embedding_model, - timeout, - }, - quantization: Quantization::Binary, - })) -} - -fn py_sys_module(py: Python<'_>) -> PyResult>> { - match py - .import("sys")? - .getattr("modules")? - .get_item("litellm.proxy.proxy_server") - { - Ok(module) => Ok(Some(module)), - Err(error) if error.is_instance_of::(py) => Ok(None), - Err(error) => Err(error), - } -} - -#[inline(never)] -fn project_azure_blob(backend: &Bound<'_, PyAny>) -> PyResult { - let client = backend.getattr("container_client")?; - let container = client.getattr("container_name")?.extract::()?; - let url = client.getattr("url")?.extract::()?; - let account_url = url - .strip_suffix(container.as_str()) - .and_then(|url| url.strip_suffix('/')) - .ok_or_else(|| PyValueError::new_err("Azure Blob container URL is malformed"))?; - Ok(AzureBlobCacheConfig { - account_url: account_url.to_string(), - container, - }) -} - -#[inline(never)] -pub(super) fn project_redis_semantic( - backend: &Bound<'_, PyAny>, -) -> PyResult { - Ok(RedisSemanticCacheConfig { - redis_url: backend.getattr("_redis_url")?.extract::()?, - index_name: backend - .getattr("_index_name")? - .extract::>()? - .unwrap_or_else(|| "litellm_semantic_cache_index".into()), - similarity_threshold: backend.getattr("similarity_threshold")?.extract::()?, - }) -} - -#[inline(never)] -fn project_memory(backend: &Bound<'_, PyAny>) -> PyResult { - let max_size_kib = backend.getattr("max_size_per_item")?.extract::()?; - Ok(MemoryCacheConfig { - default_ttl: duration(backend.getattr("default_ttl")?.extract::()?)?, - capacity: backend.getattr("max_size_in_memory")?.extract::()?, - max_entry_bytes: max_size_kib - .checked_mul(1024) - .ok_or_else(|| PyValueError::new_err("memory cache item limit is too large"))?, - }) -} - -#[inline(never)] -fn project_gcs( - backend: &Bound<'_, PyAny>, -) -> PyResult> { - let bucket_name = match backend.getattr("bucket_name")?.extract::>() { - Ok(Some(bucket_name)) if !bucket_name.is_empty() => bucket_name, - _ => return Ok(Err(UnsupportedCacheConfig::GcsBucket)), - }; - Ok(Ok(GcsCacheConfig { - bucket_name, - key_prefix: backend.getattr("key_prefix")?.extract::()?, - path_service_account: backend - .getattr("path_service_account")? - .extract::>()?, - })) -} - -#[inline(never)] -fn project_disk( - backend: &Bound<'_, PyAny>, -) -> PyResult> { - let store = backend.getattr("disk_cache")?; - if !instance_class_is(&store, "diskcache.core", "Cache")? - || !instance_class_is(&store.getattr("_disk")?, "diskcache.core", "Disk")? - { - return Ok(Err(UnsupportedCacheConfig::DiskStore)); - } - Ok(Ok(DiskCacheConfig { - directory: PathBuf::from(store.getattr("directory")?.extract::()?), - })) -} - -#[inline(never)] -fn project_redis( - backend: &Bound<'_, PyAny>, -) -> PyResult> { - let source = backend.getattr("redis_kwargs")?.cast_into::()?; - if has_value(&source, "sentinel_nodes")? { - return Ok(Err(UnsupportedCacheConfig::RedisTopology)); - } - for key in ["credential_provider", "redis_connect_func"] { - if has_value(&source, key)? { - return Ok(Err(UnsupportedCacheConfig::RedisCredentials)); - } - } - if has_value(&source, "connection_pool")? { - return Ok(Err(UnsupportedCacheConfig::RedisConnection)); - } - for key in [ - "retry", - "retry_on_error", - "socket_keepalive_options", - "unix_socket_path", - "cache", - "cache_config", - "event_dispatcher", - "ssl_ca_path", - "ssl_password", - "ssl_min_version", - "ssl_ciphers", - "ssl_validate_ocsp", - "ssl_validate_ocsp_stapled", - "ssl_ocsp_context", - "ssl_ocsp_expected_cert", - ] { - if has_value(&source, key)? { - return Ok(Err(UnsupportedCacheConfig::RedisOption)); - } - } - for key in ["retry_on_timeout", "single_connection_client"] { - if optional_coerced_bool(&source, key)?.unwrap_or(false) { - return Ok(Err(UnsupportedCacheConfig::RedisOption)); - } - } - - let client = backend.getattr("redis_client")?; - let projection = if has_value(&source, "startup_nodes")? { - project_cluster_client(&source, &client)? - } else { - project_standalone_client(&client)? - }; - let RedisClientProjection { - topology, - host, - port, - pool_size, - resolved, - tls, - } = match projection { - Ok(projection) => projection, - Err(reason) => return Ok(Err(reason)), - }; - if has_value(&resolved, "credential_provider")? { - return Ok(Err(UnsupportedCacheConfig::RedisCredentials)); - } - Ok(Ok(RedisCacheConfig { - default_ttl: duration(backend.getattr("default_ttl")?.extract::()?)?, - namespace: optional_attribute_string(backend, "namespace")?, - flush_size: backend.getattr("redis_flush_size")?.extract::()?, - topology, - connection: resolved_connection(&resolved, host, port, pool_size, tls)?, - })) -} - -/// The connection settings redis-py resolved for one client's pool. -#[inline(never)] -fn resolved_connection( - resolved: &Bound<'_, PyDict>, - host: String, - port: u16, - pool_size: usize, - tls: Option, -) -> PyResult { - let protocol = match optional_i64(resolved, "protocol")?.unwrap_or(2) { - 2 => RedisProtocol::Resp2, - 3 => RedisProtocol::Resp3, - _ => return Err(PyValueError::new_err("unsupported Redis protocol version")), - }; - Ok(RedisConnectionConfig { - host, - port, - database: optional_i64(resolved, "db")?.unwrap_or(0), - username: optional_dict_string(resolved, "username")?, - password: optional_dict_string(resolved, "password")?, - protocol, - pool_size, - read_timeout: optional_dict_duration(resolved, "socket_timeout")?, - connect_timeout: optional_dict_duration(resolved, "socket_connect_timeout")?, - socket_keepalive: optional_bool(resolved, "socket_keepalive")?, - health_check_interval: duration( - optional_f64(resolved, "health_check_interval")?.unwrap_or(0.0), - )?, - client_name: optional_dict_string(resolved, "client_name")?, - tls, - }) -} - -#[inline(never)] -fn project_s3( - backend: &Bound<'_, PyAny>, -) -> PyResult> { - let client = backend.getattr("s3_client")?; - if !instance_class_is(&client, "botocore.client", "S3")? { - return Ok(Err(UnsupportedCacheConfig::S3Client)); - } - let meta = client.getattr("meta")?; - let Some(region) = optional_string(meta.getattr("region_name")?)? else { - return Ok(Err(UnsupportedCacheConfig::S3Option)); - }; - let Some(endpoint_url) = optional_string(meta.getattr("endpoint_url")?)? else { - return Ok(Err(UnsupportedCacheConfig::S3Option)); - }; - let client_config = meta.getattr("config")?; - for name in ["s3", "proxies", "client_cert"] { - if optional_attribute(&client_config, name)?.is_some_and(|value| !value.is_none()) { - return Ok(Err(UnsupportedCacheConfig::S3Option)); - } - } - let signature = match optional_attribute(&client_config, "signature_version")? { - Some(value) => value.extract::>()?, - None => None, - }; - if signature.as_deref() != Some("s3v4") { - return Ok(Err(UnsupportedCacheConfig::S3Option)); - } - let insecure = endpoint_url.starts_with("http://"); - let verify = optional_attribute_chain(&client, &["_endpoint", "http_session", "_verify"])?; - let verified = verify - .and_then(|value| value.cast::().ok().map(|value| value.is_true())) - .unwrap_or(false); - if !verified && !insecure { - return Ok(Err(UnsupportedCacheConfig::S3Option)); - } - let credentials = optional_attribute_chain(&client, &["_request_signer", "_credentials"])? - .ok_or(UnsupportedCacheConfig::S3Credentials); - let credentials = match credentials { - Ok(credentials) if !credentials.is_none() => credentials, - _ => return Ok(Err(UnsupportedCacheConfig::S3Credentials)), - }; - let auth = if credentials.getattr("method")?.extract::()?.as_str() == "explicit" { - AwsAuthConfig { - access_key_id: credentials - .getattr("access_key")? - .extract::>()?, - secret_access_key: credentials - .getattr("secret_key")? - .extract::>()?, - session_token: credentials.getattr("token")?.extract::>()?, - region_name: Some(region.clone()), - ..Default::default() - } - } else { - AwsAuthConfig { - region_name: Some(region.clone()), - ..Default::default() - } - }; - let default_endpoint = endpoint_url == format!("https://s3.{region}.amazonaws.com") - || (region == "us-east-1" && endpoint_url == "https://s3.amazonaws.com"); - Ok(Ok(S3CacheConfig { - bucket: backend.getattr("bucket_name")?.extract::()?, - key_prefix: backend.getattr("key_prefix")?.extract::()?, - region, - endpoint: (!default_endpoint).then_some(S3Endpoint { url: endpoint_url }), - auth, - })) -} - -#[inline(never)] -fn project_standalone_client<'py>( - client: &Bound<'py, PyAny>, -) -> PyResult, UnsupportedCacheConfig>> { - let pool = client.getattr("connection_pool")?; - if !instance_class_is(&pool, "redis.connection", "ConnectionPool")? { - return Ok(Err(UnsupportedCacheConfig::RedisConnection)); - } - let resolved = pool.getattr("connection_kwargs")?.cast_into::()?; - if has_value(&resolved, "redis_connect_func")? { - return Ok(Err(UnsupportedCacheConfig::RedisCredentials)); - } - let connection_class = resolved - .get_item("connection_class")? - .unwrap_or(pool.getattr("connection_class")?); - let tls = if class_is(&connection_class, "redis.connection", "Connection")? { - None - } else if class_is(&connection_class, "redis.connection", "SSLConnection")? { - Some(project_tls(&resolved)?) - } else { - return Ok(Err(UnsupportedCacheConfig::RedisConnection)); - }; - Ok(Ok(RedisClientProjection { - topology: RedisTopology::Standalone, - host: required_string(&resolved, "host")?, - port: port(required_i64(&resolved, "port")?)?, - pool_size: pool.getattr("max_connections")?.extract::()?, - resolved, - tls, - })) -} - -#[inline(never)] -fn project_cluster_client<'py>( - source: &Bound<'_, PyDict>, - client: &Bound<'py, PyAny>, -) -> PyResult, UnsupportedCacheConfig>> { - let Some(startup_nodes) = startup_nodes(source)? else { - return Ok(Err(UnsupportedCacheConfig::RedisTopology)); - }; - if !instance_class_is(client, "redis.cluster", "RedisCluster")? { - return Ok(Err(UnsupportedCacheConfig::RedisConnection)); - } - let nodes = client.getattr("nodes_manager")?; - if !class_is( - &nodes.getattr("connection_pool_class")?, - "redis.connection", - "ConnectionPool", - )? { - return Ok(Err(UnsupportedCacheConfig::RedisConnection)); - } - let resolved = nodes.getattr("connection_kwargs")?.cast_into::()?; - if let Some(connect) = resolved.get_item("redis_connect_func")? - && !connect.is_none() - { - let own_hook = connect - .getattr("__self__") - .is_ok_and(|owner| owner.is(client)) - && connect - .getattr("__func__") - .and_then(|function| Ok(function.is(&client.get_type().getattr("on_connect")?))) - .unwrap_or(false); - if !own_hook { - return Ok(Err(UnsupportedCacheConfig::RedisCredentials)); - } - } - let tls = if optional_bool(&resolved, "ssl")?.unwrap_or(false) { - Some(project_tls(&resolved)?) - } else { - None - }; - let first = &startup_nodes[0]; - Ok(Ok(RedisClientProjection { - host: first.host.clone(), - port: first.port, - pool_size: optional_i64(&resolved, "max_connections")? - .map(|value| { - usize::try_from(value).map_err(|_| PyValueError::new_err("invalid Redis pool size")) - }) - .transpose()? - .unwrap_or(REDIS_PY_DEFAULT_MAX_CONNECTIONS), - topology: RedisTopology::Cluster { startup_nodes }, - resolved, - tls, - })) -} - -#[inline(never)] -fn startup_nodes(source: &Bound<'_, PyDict>) -> PyResult>> { - let Some(nodes) = source.get_item("startup_nodes")? else { - return Ok(None); - }; - let Ok(nodes) = nodes.cast_into::() else { - return Ok(None); - }; - if nodes.is_empty() { - return Ok(None); - } - let mut parsed = Vec::with_capacity(nodes.len()); - for node in nodes.iter() { - let Ok(node) = node.cast_into::() else { - return Ok(None); - }; - if node.len() != 2 || !has_value(&node, "host")? || !has_value(&node, "port")? { - return Ok(None); - } - let (Ok(host), Ok(port)) = ( - required_string(&node, "host"), - required_i64(&node, "port").and_then(port), - ) else { - return Ok(None); - }; - parsed.push(RedisNode { host, port }); - } - Ok(Some(parsed)) -} - -#[inline(never)] -fn port(value: i64) -> PyResult { - u16::try_from(value).map_err(|_| PyValueError::new_err("invalid Redis port")) -} - -#[inline(never)] -fn project_valkey_semantic( - backend: &Bound<'_, PyAny>, -) -> PyResult> { - let client = backend.getattr("sync_client")?; - let pool = client.getattr("connection_pool")?; - let Ok((resolved, is_tls)) = project_connection_pool(&pool)? else { - return Ok(Err(UnsupportedCacheConfig::RedisConnection)); - }; - for key in ["credential_provider", "redis_connect_func"] { - if has_value(&resolved, key)? { - return Ok(Err(UnsupportedCacheConfig::RedisCredentials)); - } - } - if is_tls { - return Ok(Err(UnsupportedCacheConfig::ValkeyTls)); - } - let host = required_string(&resolved, "host")?; - if host.is_empty() { - return Ok(Err(UnsupportedCacheConfig::RedisConnection)); - } - let connection = resolved_connection( - &resolved, - host, - port(required_i64(&resolved, "port")?)?, - pool.getattr("max_connections")?.extract::()?, - None, - )?; - Ok(Ok(ValkeySemanticCacheConfig { - similarity_threshold: backend.getattr("similarity_threshold")?.extract()?, - index_name: backend.getattr("index_name")?.extract()?, - connection, - })) -} - -#[inline(never)] -fn project_connection_pool<'py>( - pool: &Bound<'py, PyAny>, -) -> PyResult, bool), UnsupportedCacheConfig>> { - if !instance_class_is(pool, "redis.connection", "ConnectionPool")? { - return Ok(Err(UnsupportedCacheConfig::RedisConnection)); - } - let resolved = pool.getattr("connection_kwargs")?.cast_into::()?; - let connection_class = resolved - .get_item("connection_class")? - .unwrap_or(pool.getattr("connection_class")?); - let is_tls = if class_is(&connection_class, "redis.connection", "Connection")? { - false - } else if class_is(&connection_class, "redis.connection", "SSLConnection")? { - true - } else { - return Ok(Err(UnsupportedCacheConfig::RedisConnection)); - }; - Ok(Ok((resolved, is_tls))) -} - -#[inline(never)] -fn project_tls(values: &Bound<'_, PyDict>) -> PyResult { - Ok(RedisTlsConfig { - certificate_requirement: certificate_requirement(values)?, - check_hostname: optional_bool(values, "ssl_check_hostname")?.unwrap_or(false), - ca_certificate: optional_dict_string(values, "ssl_ca_certs")?, - ca_data: optional_dict_string(values, "ssl_ca_data")?, - client_certificate: optional_dict_string(values, "ssl_certfile")?, - client_key: optional_dict_string(values, "ssl_keyfile")?, - }) -} - -#[inline(never)] -fn certificate_requirement(values: &Bound<'_, PyDict>) -> PyResult { - let Some(value) = values.get_item("ssl_cert_reqs")? else { - return Ok(CertificateRequirement::Required); - }; - if value.is_none() { - return Ok(CertificateRequirement::Required); - } - if let Ok(number) = value.extract::() { - return match number { - 0 => Ok(CertificateRequirement::None), - 1 => Ok(CertificateRequirement::Optional), - 2 => Ok(CertificateRequirement::Required), - _ => Err(PyValueError::new_err( - "invalid Redis TLS certificate requirement", - )), - }; - } - let text = value.str()?; - let text = text.to_str()?; - if text.eq_ignore_ascii_case("none") || text.eq_ignore_ascii_case("cert_none") { - return Ok(CertificateRequirement::None); - } - if text.eq_ignore_ascii_case("optional") || text.eq_ignore_ascii_case("cert_optional") { - return Ok(CertificateRequirement::Optional); - } - if text.eq_ignore_ascii_case("required") || text.eq_ignore_ascii_case("cert_required") { - return Ok(CertificateRequirement::Required); - } - Err(PyValueError::new_err( - "invalid Redis TLS certificate requirement", - )) -} - -#[inline(never)] -fn instance_class_is(value: &Bound<'_, PyAny>, module: &str, name: &str) -> PyResult { - class_is(value.get_type().as_any(), module, name) -} - -#[inline(never)] -fn class_is(value: &Bound<'_, PyAny>, module: &str, name: &str) -> PyResult { - Ok(value - .getattr("__module__")? - .cast_into::()? - .to_str()? - == module - && value - .getattr("__qualname__")? - .cast_into::()? - .to_str()? - == name) -} - -#[inline(never)] -fn optional_attribute_string(value: &Bound<'_, PyAny>, name: &str) -> PyResult> { - match value.getattr(name) { - Ok(value) => optional_string(value), - Err(error) if error.is_instance_of::(value.py()) => { - Ok(None) - } - Err(error) => Err(error), - } -} - -#[inline(never)] -fn optional_attribute<'py>( - value: &Bound<'py, PyAny>, - name: &str, -) -> PyResult>> { - match value.getattr(name) { - Ok(value) => Ok(Some(value)), - Err(error) if error.is_instance_of::(value.py()) => Ok(None), - Err(error) => Err(error), - } -} - -#[inline(never)] -fn optional_attribute_chain<'py>( - value: &Bound<'py, PyAny>, - names: &[&str], -) -> PyResult>> { - names - .iter() - .try_fold(Some(value.clone()), |current, name| match current { - Some(current) => optional_attribute(¤t, name), - None => Ok(None), - }) -} - -#[inline(never)] -fn optional_string(value: Bound<'_, PyAny>) -> PyResult> { - Ok(value - .extract::>()? - .filter(|value| !value.is_empty())) -} - -#[inline(never)] -fn has_value(values: &Bound<'_, PyDict>, key: &str) -> PyResult { - Ok(values.get_item(key)?.is_some_and(|value| !value.is_none())) -} - -#[inline(never)] -fn required_string(values: &Bound<'_, PyDict>, key: &str) -> PyResult { - values - .get_item(key)? - .ok_or_else(|| PyTypeError::new_err("Redis connection is incomplete"))? - .extract::() -} - -#[inline(never)] -fn required_i64(values: &Bound<'_, PyDict>, key: &str) -> PyResult { - values - .get_item(key)? - .ok_or_else(|| PyTypeError::new_err("Redis connection is incomplete"))? - .extract::() -} - -#[inline(never)] -fn optional_dict_string(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { - match values.get_item(key)? { - Some(value) if !value.is_none() => optional_string(value), - _ => Ok(None), - } -} - -#[inline(never)] -fn optional_f64(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { - match values.get_item(key)? { - Some(value) => value.extract::>(), - None => Ok(None), - } -} - -#[inline(never)] -fn optional_i64(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { - match values.get_item(key)? { - Some(value) => value.extract::>(), - None => Ok(None), - } -} - -#[inline(never)] -fn optional_bool(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { - match values.get_item(key)? { - Some(value) => value.extract::>(), - None => Ok(None), - } -} - -#[inline(never)] -fn optional_coerced_bool(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { - let Some(value) = values.get_item(key)? else { - return Ok(None); - }; - if value.is_none() { - return Ok(None); - } - if let Ok(text) = value.extract::() { - return Ok(Some( - text == "1" || text.eq_ignore_ascii_case("true") || text.eq_ignore_ascii_case("yes"), - )); - } - value.extract::().map(Some) -} - -#[inline(never)] -fn optional_dict_duration(values: &Bound<'_, PyDict>, key: &str) -> PyResult> { - optional_f64(values, key)?.map(duration).transpose() -} - -#[cfg(test)] -mod tests { - use std::{ffi::CString, time::Duration}; - - use litellm_cache_redis::{RedisNode, RedisTopology}; - use litellm_cache_redis_semantic::RedisSemanticConfig; - use pyo3::{prelude::*, types::PyDict}; - use rstest::{fixture, rstest}; - - use super::{ - CacheBackendConfig, CacheConfigProjection, CachePolicy, CertificateRequirement, - GcsCacheConfig, NativeCacheConfig, REDIS_PY_DEFAULT_MAX_CONNECTIONS, RedisConnectionConfig, - RedisProtocol, RedisSemanticCacheConfig, RedisTlsConfig, UnsupportedCacheConfig, - }; - use crate::cache::native::{backend::NativeResponseCache, embedder::PythonEmbedder}; - - fn cluster_facade<'py>(py: Python<'py>, startup_nodes: &str, hook: &str) -> Bound<'py, PyAny> { - facade( - py, - &format!( - "RedisCluster = type('RedisCluster', (), {{'__module__': 'redis.cluster', 'on_connect': lambda self, connection: None}})\n\ - client = RedisCluster()\n\ - client.nodes_manager = SimpleNamespace(connection_pool_class=ConnectionPool, connection_kwargs={{'password': 'secret', 'redis_connect_func': {hook}, 'protocol': 3, 'ssl': True, 'ssl_cert_reqs': 'none'}})\n\ - backend = SimpleNamespace(default_ttl=120, namespace='team', redis_flush_size=100, redis_kwargs={{'startup_nodes': {startup_nodes}, 'password': 'secret'}}, redis_client=client)\n\ - facade = SimpleNamespace(type='redis', mode='default-on', ttl=None, namespace='team', supported_call_types=None, redis_flush_size=100, semantic_cache_scope='key', cache=backend)" - ), - ) - } - - fn facade<'py>(py: Python<'py>, body: &str) -> Bound<'py, PyAny> { - let locals = PyDict::new(py); - py.run( - &CString::new(format!( - "from types import SimpleNamespace\n\ - ConnectionPool = type('ConnectionPool', (), {{'__module__': 'redis.connection'}})\n\ - Connection = type('Connection', (), {{'__module__': 'redis.connection'}})\n\ - SSLConnection = type('SSLConnection', (), {{'__module__': 'redis.connection'}})\n\ - {body}" - )) - .unwrap(), - None, - Some(&locals), - ) - .unwrap(); - locals.get_item("facade").unwrap().unwrap() - } - - fn native(facade: &Bound<'_, PyAny>) -> NativeCacheConfig { - match NativeCacheConfig::project(facade).unwrap() { - CacheConfigProjection::Native(config) => *config, - CacheConfigProjection::Unsupported(reason) => panic!("{}", reason.message()), - } - } - - fn unsupported(facade: &Bound<'_, PyAny>) -> UnsupportedCacheConfig { - match NativeCacheConfig::project(facade).unwrap() { - CacheConfigProjection::Native(_) => panic!("configuration must stay on Python"), - CacheConfigProjection::Unsupported(reason) => reason, - } - } - - #[fixture] - fn interpreter() { - Python::initialize(); - } - - #[fixture] - fn connection() -> RedisConnectionConfig { - RedisConnectionConfig { - host: "cache.internal".into(), - port: 6380, - database: 4, - username: None, - password: None, - protocol: RedisProtocol::Resp2, - pool_size: REDIS_PY_DEFAULT_MAX_CONNECTIONS, - read_timeout: None, - connect_timeout: None, - socket_keepalive: None, - health_check_interval: Duration::ZERO, - client_name: None, - tls: None, - } - } - - fn verified_tls() -> RedisTlsConfig { - RedisTlsConfig { - certificate_requirement: CertificateRequirement::Required, - check_hostname: true, - ca_certificate: None, - ca_data: None, - client_certificate: None, - client_key: None, - } - } - - #[rstest] - fn projects_effective_memory_configuration(_interpreter: ()) { - Python::attach(|py| { - let facade = facade( - py, - "backend = SimpleNamespace(default_ttl=913, max_size_in_memory=37, max_size_per_item=8)\n\ - facade = SimpleNamespace(type='local', mode='default-on', ttl=11.5, namespace=None, supported_call_types=['completion'], redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - ); - let config = native(&facade); - assert_eq!(config.policy.semantic_cache_scope, "key"); - assert_eq!(config.policy.redis_flush_size, None); - let CacheBackendConfig::Memory(memory) = config.backend else { - panic!("expected memory configuration"); - }; - assert_eq!(memory.default_ttl, Duration::from_secs(913)); - assert_eq!(memory.capacity, 37); - assert_eq!(memory.max_entry_bytes, 8192); - let matching = NativeResponseCache::memory(37, Duration::from_secs(913), 8192); - let mismatched = NativeResponseCache::memory(37, Duration::from_secs(913), 8191); - let matching_config = NativeCacheConfig { - policy: config.policy, - backend: CacheBackendConfig::Memory(memory), - }; - assert_eq!(matching_config.service_mismatch(&matching), None); - assert_eq!( - matching_config.service_mismatch(&mismatched), - Some("facade and native backend item limits must match") - ); - }); - } - - #[rstest] - fn redis_semantic_service_mismatch_accepts_backend_precision_threshold(_interpreter: ()) { - Python::attach(|py| { - let facade = facade( - py, - "backend = SimpleNamespace(_redis_url='redis://127.0.0.1/', _index_name='semantic_idx', similarity_threshold=0.8, embedding_model='text-embedding-3-small', embedding_max_input_tokens=None, embedding_timeout=None)\n\ - facade = SimpleNamespace(type='redis-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - ); - let backend = facade.getattr("cache").unwrap(); - let embedder = PythonEmbedder::new(backend.clone().unbind()); - let CacheBackendConfig::RedisSemantic(config) = native(&facade).backend else { - panic!("expected Redis semantic configuration"); - }; - let service = NativeResponseCache::redis_semantic( - &config.redis_url, - embedder, - RedisSemanticConfig { - index_name: config.index_name.clone(), - similarity_threshold: config.similarity_threshold as f32, - }, - ) - .unwrap(); - let matching_config = NativeCacheConfig { - policy: CachePolicy { - redis_flush_size: None, - semantic_cache_scope: "key".into(), - }, - backend: CacheBackendConfig::RedisSemantic(config), - }; - assert_eq!(matching_config.service_mismatch(&service), None); - }); - } - - #[rstest] - fn projects_resolved_redis_tls_configuration(_interpreter: ()) { - Python::attach(|py| { - let facade = facade( - py, - "pool = ConnectionPool()\n\ - pool.connection_class = SSLConnection\n\ - pool.max_connections = 29\n\ - pool.connection_kwargs = {'host': 'cache.internal', 'port': 6380, 'db': 4, 'username': 'user', 'password': 'secret', 'protocol': 3, 'socket_timeout': 7.5, 'socket_connect_timeout': 2, 'socket_keepalive': True, 'health_check_interval': 15, 'client_name': 'litellm', 'ssl_cert_reqs': 'optional', 'ssl_check_hostname': True, 'ssl_ca_certs': '/ca.pem', 'ssl_ca_data': 'CA DATA', 'ssl_certfile': '/client.pem', 'ssl_keyfile': '/client.key'}\n\ - client = SimpleNamespace(connection_pool=pool)\n\ - backend = SimpleNamespace(default_ttl=777, namespace='team', redis_flush_size=31, redis_kwargs={}, redis_client=client)\n\ - facade = SimpleNamespace(type='redis', mode='default-off', ttl=None, namespace='team', supported_call_types=None, redis_flush_size=31, semantic_cache_scope='key', cache=backend)", - ); - let config = native(&facade); - assert_eq!(config.policy.redis_flush_size, Some(31)); - let CacheBackendConfig::Redis(redis) = config.backend else { - panic!("expected Redis configuration"); - }; - assert_eq!(redis.default_ttl, Duration::from_secs(777)); - assert_eq!(redis.namespace.as_deref(), Some("team")); - assert_eq!(redis.flush_size, 31); - assert_eq!(redis.connection.host, "cache.internal"); - assert_eq!(redis.connection.port, 6380); - assert_eq!(redis.connection.database, 4); - assert_eq!(redis.connection.protocol, RedisProtocol::Resp3); - assert_eq!(redis.connection.pool_size, 29); - assert_eq!( - redis.connection.read_timeout, - Some(Duration::from_secs_f64(7.5)) - ); - assert_eq!( - redis.connection.connect_timeout, - Some(Duration::from_secs(2)) - ); - assert_eq!(redis.connection.socket_keepalive, Some(true)); - assert_eq!( - redis.connection.health_check_interval, - Duration::from_secs(15) - ); - assert_eq!(redis.connection.client_name.as_deref(), Some("litellm")); - let tls = redis.connection.tls.as_ref().unwrap(); - assert_eq!( - tls.certificate_requirement, - CertificateRequirement::Optional - ); - assert!(tls.check_hostname); - assert_eq!(tls.ca_certificate.as_deref(), Some("/ca.pem")); - assert_eq!(tls.ca_data.as_deref(), Some("CA DATA")); - assert_eq!(tls.client_certificate.as_deref(), Some("/client.pem")); - assert_eq!(tls.client_key.as_deref(), Some("/client.key")); - assert!(matches!( - redis.connection.native_url(), - Err(UnsupportedCacheConfig::RedisPoolSize) - )); - }); - } - - #[rstest] - fn projects_valkey_semantic_configuration(_interpreter: ()) { - Python::attach(|py| { - let facade = facade( - py, - "pool = ConnectionPool()\n\ - pool.connection_class = Connection\n\ - pool.max_connections = 12\n\ - pool.connection_kwargs = {'host': 'cache.internal', 'port': 6390, 'db': 2, 'socket_timeout': 3}\n\ - client = SimpleNamespace(connection_pool=pool)\n\ - backend = SimpleNamespace(similarity_threshold=0.85, index_name='semantic_idx', embedding_model='text-embedding-3-small', sync_client=client)\n\ - facade = SimpleNamespace(type='valkey-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - ); - let CacheBackendConfig::ValkeySemantic(valkey) = native(&facade).backend else { - panic!("expected Valkey semantic configuration"); - }; - assert_eq!(valkey.similarity_threshold, 0.85); - assert_eq!(valkey.index_name, "semantic_idx"); - assert_eq!(valkey.connection.host, "cache.internal"); - assert_eq!(valkey.connection.port, 6390); - assert_eq!(valkey.connection.database, 2); - assert_eq!(valkey.connection.pool_size, 12); - assert_eq!(valkey.connection.protocol, RedisProtocol::Resp2); - assert_eq!(valkey.connection.read_timeout, Some(Duration::from_secs(3))); - assert!(valkey.connection.tls.is_none()); - }); - } - - #[rstest] - #[case::valkey_tls( - "pool = ConnectionPool()\n\ - pool.connection_class = SSLConnection\n\ - pool.connection_kwargs = {'host': 'cache.internal', 'port': 6390}\n\ - client = SimpleNamespace(connection_pool=pool)\n\ - backend = SimpleNamespace(similarity_threshold=0.85, index_name='semantic_idx', embedding_model='text-embedding-3-small', sync_client=client)\n\ - facade = SimpleNamespace(type='valkey-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - "native Valkey semantic cache does not support TLS connections" - )] - #[case::valkey_dynamic_auth( - "pool = ConnectionPool()\n\ - pool.connection_class = Connection\n\ - pool.connection_kwargs = {'host': 'cache.internal', 'port': 6390, 'credential_provider': object()}\n\ - client = SimpleNamespace(connection_pool=pool)\n\ - backend = SimpleNamespace(similarity_threshold=0.85, index_name='semantic_idx', embedding_model='text-embedding-3-small', sync_client=client)\n\ - facade = SimpleNamespace(type='valkey-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - "native Redis credentials require Python" - )] - #[case::redis_dynamic_auth( - "backend = SimpleNamespace(redis_kwargs={'credential_provider': object()})\n\ - facade = SimpleNamespace(type='redis', mode='default-on', ttl=None, namespace=None, supported_call_types=[], redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - "native Redis credentials require Python" - )] - #[case::gcs_without_bucket( - "backend = SimpleNamespace(bucket_name=None, key_prefix='', path_service_account=None)\n\ - facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - "native GCS cache requires a configured bucket name" - )] - fn configurations_that_stay_on_python( - _interpreter: (), - #[case] body: &str, - #[case] message: &str, - ) { - Python::attach(|py| { - assert_eq!(unsupported(&facade(py, body)).message(), message); - }); - } - - #[rstest] - fn projects_cluster_startup_nodes_as_redis_topology(_interpreter: ()) { - Python::attach(|py| { - let facade = cluster_facade( - py, - "[{'host': 'node-a', 'port': 7000}, {'host': 'node-b', 'port': 7001}]", - "client.on_connect", - ); - let CacheBackendConfig::Redis(redis) = native(&facade).backend else { - panic!("expected Redis configuration"); - }; - let expected = RedisTopology::Cluster { - startup_nodes: vec![ - RedisNode { - host: "node-a".into(), - port: 7000, - }, - RedisNode { - host: "node-b".into(), - port: 7001, - }, - ], - }; - assert_eq!(redis.topology, expected); - assert_eq!(redis.connection.host, "node-a"); - assert_eq!(redis.connection.port, 7000); - assert_eq!(redis.connection.password.as_deref(), Some("secret")); - assert_eq!(redis.connection.protocol, RedisProtocol::Resp3); - assert_eq!( - redis - .connection - .tls - .as_ref() - .unwrap() - .certificate_requirement, - CertificateRequirement::None - ); - assert!(matches!( - redis.connection.native_url(), - Err(UnsupportedCacheConfig::RedisTlsVerification) - )); - }); - } - - #[rstest] - fn projects_gcs_configuration(_interpreter: ()) { - Python::attach(|py| { - let facade = facade( - py, - "backend = SimpleNamespace(bucket_name='bucket', key_prefix='cache/', path_service_account='credentials.json')\n\ - facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - ); - let config = native(&facade); - let CacheBackendConfig::Gcs(gcs) = config.backend else { - panic!("expected GCS configuration"); - }; - assert_eq!( - gcs, - GcsCacheConfig { - bucket_name: "bucket".into(), - key_prefix: "cache/".into(), - path_service_account: Some("credentials.json".into()), - } - ); - let matching = NativeResponseCache::gcs( - litellm_cache_gcs::GcsConfig { - bucket_name: "bucket".into(), - gcs_path: Some("cache/".into()), - path_service_account: Some("credentials.json".into()), - endpoint: litellm_cache_gcs::DEFAULT_ENDPOINT.into(), - }, - litellm_http::Client::plain_for_test(), - Some("token".into()), - ); - let matching_config = NativeCacheConfig { - policy: config.policy, - backend: CacheBackendConfig::Gcs(gcs), - }; - assert_eq!(matching_config.service_mismatch(&matching), None); - }); - } - - #[rstest] - #[case::extra_node_field( - "[{'host': 'node-a', 'port': 7000, 'server_type': 'primary'}]", - "client.on_connect", - "native Redis topology is not implemented" - )] - #[case::non_numeric_port( - "[{'host': 'node-a', 'port': 'seven'}]", - "client.on_connect", - "native Redis topology is not implemented" - )] - #[case::empty("[]", "client.on_connect", "native Redis topology is not implemented")] - #[case::foreign_hook( - "[{'host': 'node-a', 'port': 7000}]", - "lambda connection: None", - "native Redis credentials require Python" - )] - fn malformed_startup_nodes_and_foreign_connect_hooks_stay_on_python( - _interpreter: (), - #[case] startup_nodes: &str, - #[case] hook: &str, - #[case] message: &str, - ) { - Python::attach(|py| { - let facade = cluster_facade(py, startup_nodes, hook); - assert_eq!(unsupported(&facade).message(), message); - }); - } - - #[rstest] - #[case::plain(|_: &mut RedisConnectionConfig| {}, "redis://cache.internal:6380/4")] - #[case::credentials( - |connection: &mut RedisConnectionConfig| { - connection.username = Some("user".into()); - connection.password = Some("p@ss:word".into()); - }, - "redis://user:p%40ss%3Aword@cache.internal:6380/4" - )] - #[case::password_only( - |connection: &mut RedisConnectionConfig| connection.password = Some("secret".into()), - "redis://:secret@cache.internal:6380/4" - )] - #[case::resp3( - |connection: &mut RedisConnectionConfig| connection.protocol = RedisProtocol::Resp3, - "redis://cache.internal:6380/4?protocol=resp3" - )] - #[case::ipv6( - |connection: &mut RedisConnectionConfig| connection.host = "::1".into(), - "redis://[::1]:6380/4" - )] - #[case::verified_tls( - |connection: &mut RedisConnectionConfig| connection.tls = Some(verified_tls()), - "rediss://cache.internal:6380/4" - )] - #[case::optional_certificate( - |connection: &mut RedisConnectionConfig| { - connection.tls = Some(RedisTlsConfig { - certificate_requirement: CertificateRequirement::Optional, - ..verified_tls() - }); - }, - "rediss://cache.internal:6380/4" - )] - #[case::keepalive_off( - |connection: &mut RedisConnectionConfig| connection.socket_keepalive = Some(false), - "redis://cache.internal:6380/4" - )] - #[case::redis_cache_socket_timeout( - |connection: &mut RedisConnectionConfig| { - connection.read_timeout = Some(Duration::from_secs(5)); - }, - "redis://cache.internal:6380/4" - )] - fn native_url_encodes_the_resolved_connection( - mut connection: RedisConnectionConfig, - #[case] configure: fn(&mut RedisConnectionConfig), - #[case] expected: &str, - ) { - configure(&mut connection); - assert_eq!(connection.native_url().ok().as_deref(), Some(expected)); - } - - #[rstest] - #[case::pool_size( - |connection: &mut RedisConnectionConfig| connection.pool_size = 50, - "native Redis uses a fixed connection pool; max_connections requires Python" - )] - #[case::socket_timeout( - |connection: &mut RedisConnectionConfig| { - connection.read_timeout = Some(Duration::from_millis(100)); - }, - "native Redis uses fixed socket timeouts; socket_timeout and socket_connect_timeout require Python" - )] - #[case::connect_timeout( - |connection: &mut RedisConnectionConfig| { - connection.connect_timeout = Some(Duration::from_secs(1)); - }, - "native Redis uses fixed socket timeouts; socket_timeout and socket_connect_timeout require Python" - )] - #[case::keepalive( - |connection: &mut RedisConnectionConfig| connection.socket_keepalive = Some(true), - "native Redis does not support socket_keepalive" - )] - #[case::health_check( - |connection: &mut RedisConnectionConfig| { - connection.health_check_interval = Duration::from_secs(25); - }, - "native Redis does not support health_check_interval" - )] - #[case::client_name( - |connection: &mut RedisConnectionConfig| connection.client_name = Some("litellm".into()), - "native Redis does not support client_name" - )] - #[case::custom_ca( - |connection: &mut RedisConnectionConfig| { - connection.tls = Some(RedisTlsConfig { - ca_certificate: Some("/ca.pem".into()), - ..verified_tls() - }); - }, - "native Redis TLS does not support ssl_ca_certs, ssl_ca_data, ssl_certfile or ssl_keyfile" - )] - #[case::client_certificate( - |connection: &mut RedisConnectionConfig| { - connection.tls = Some(RedisTlsConfig { - client_certificate: Some("/client.pem".into()), - client_key: Some("/client.key".into()), - ..verified_tls() - }); - }, - "native Redis TLS does not support ssl_ca_certs, ssl_ca_data, ssl_certfile or ssl_keyfile" - )] - #[case::unverified( - |connection: &mut RedisConnectionConfig| { - connection.tls = Some(RedisTlsConfig { - certificate_requirement: CertificateRequirement::None, - check_hostname: false, - ..verified_tls() - }); - }, - "native Redis TLS always verifies the certificate and hostname; ssl_cert_reqs=none and ssl_check_hostname=false require Python" - )] - #[case::hostname_unchecked( - |connection: &mut RedisConnectionConfig| { - connection.tls = Some(RedisTlsConfig { - check_hostname: false, - ..verified_tls() - }); - }, - "native Redis TLS always verifies the certificate and hostname; ssl_cert_reqs=none and ssl_check_hostname=false require Python" - )] - fn native_url_declines_settings_the_native_client_cannot_honor( - mut connection: RedisConnectionConfig, - #[case] configure: fn(&mut RedisConnectionConfig), - #[case] message: &str, - ) { - configure(&mut connection); - let Err(reason) = connection.native_url() else { - panic!("{message}"); - }; - assert_eq!(reason.message(), message); - } - - #[rstest] - #[case::plain("redis://:secret@127.0.0.1:6379", true)] - #[case::database("redis://127.0.0.1:6379/2", true)] - #[case::unix("unix:///tmp/redis.sock", true)] - #[case::tls("rediss://cache.internal:6380", false)] - #[case::query_options("redis://127.0.0.1:6379?socket_timeout=1", false)] - #[case::malformed("not a url", false)] - fn redis_semantic_native_url_accepts_only_plain_urls(#[case] url: &str, #[case] native: bool) { - let config = RedisSemanticCacheConfig { - redis_url: url.into(), - index_name: "idx".into(), - similarity_threshold: 0.8, - }; - match config.native_url() { - Ok(value) => { - assert!(native, "{url} must decline"); - assert_eq!(value, url); - } - Err(reason) => { - assert!(!native, "{url} must be native"); - assert_eq!( - reason.message(), - "native Redis semantic cache does not support TLS or query options in redis_url" - ); - } - } - } -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native/embedder.rs b/litellm-rust/crates/python-bridge/src/cache/native/embedder.rs deleted file mode 100644 index cce36531efd..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/native/embedder.rs +++ /dev/null @@ -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, 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( - vector: Result, Error>, - future: F, -) -> impl Future { - PREPARED_EMBEDDING.scope(vector, future) -} - -/// The Python object that owns embedding for a semantic backend. -pub(in crate::cache) struct PythonEmbedder(Py); - -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) -> Self { - Self(object) - } - - pub(in crate::cache) fn object(&self) -> &Py { - &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> { - 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> { - 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> { - Ok(vector - .extract::>()? - .into_iter() - .map(|value| value as f32) - .collect()) - } - - fn embed_sync(&self, prompt: &str, metadata: Option<&Value>) -> Result, 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, 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, Error> { - self.embed_sync(prompt, metadata) - } - - fn async_embed( - &self, - _prompt: &str, - _metadata: Option<&Value>, - ) -> impl Future, 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])); - } -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native/facade.rs b/litellm-rust/crates/python-bridge/src/cache/native/facade.rs deleted file mode 100644 index 0402381777c..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/native/facade.rs +++ /dev/null @@ -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, - attributes: Vec<(String, Py)>, -} - -struct ObjectGuard { - reference: Py, - classes: Vec, - config_names: &'static [&'static str], - config: Vec, -} - -struct RedisPoolGuard { - reference: Py, - connection_class: Py, - connection_kwargs: Py, - max_connections: Option, - client_name: &'static str, - attributes: RedisPoolAttributes, -} - -struct DiskStoreGuard { - reference: Py, - directory: String, -} - -struct AzureBlobClientGuard { - sync_client: Py, - async_client: Py, - url: String, - container_name: String, -} - -struct S3ClientGuard { - reference: Py, -} - -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, - connection: ConnectionGuard, -} - -impl ObjectGuard { - fn capture( - py: Python<'_>, - object: &Bound<'_, PyAny>, - config_names: &'static [&'static str], - ) -> PyResult { - let classes = object - .get_type() - .getattr("__mro__")? - .cast_into::()? - .iter() - .map(|class| { - let class = class.cast_into::()?; - let attributes = class - .getattr("__dict__")? - .call_method0("items")? - .try_iter()? - .map(|item| item?.extract::<(String, Py)>()) - .collect::>>()?; - Ok(ClassGuard { - class: class.unbind(), - attributes, - }) - }) - .collect::>>()?; - 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> { - names - .iter() - .map(|name| match object.getattr(*name) { - Ok(value) => from_py(&value), - Err(error) - if error.is_instance_of::(object.py()) => - { - Ok(Value::Null) - } - Err(error) => Err(error), - }) - .collect() - } - - fn matches(&self, py: Python<'_>, object: &Bound<'_, PyAny>) -> PyResult { - if !self.reference.bind(py).call0()?.is(object) { - return Ok(false); - } - let mro = object - .get_type() - .getattr("__mro__")? - .cast_into::()?; - if mro.len() != self.classes.len() { - return Ok(false); - } - let instance = object.getattr("__dict__")?.cast_into::()?; - 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 { - 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::()) - .transpose()?, - client_name, - attributes, - }) - } - - fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult { - 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::()) - .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 { - 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 { - let store = backend.getattr("disk_cache")?; - Ok(self.reference.bind(py).is(&store) - && self.directory == store.getattr("directory")?.extract::()?) - } - - fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.reference) - } -} - -impl AzureBlobClientGuard { - fn capture(backend: &Bound<'_, PyAny>) -> PyResult { - let sync_client = backend.getattr("container_client")?; - Ok(Self { - url: sync_client.getattr("url")?.extract::()?, - container_name: sync_client.getattr("container_name")?.extract::()?, - sync_client: sync_client.unbind(), - async_client: backend.getattr("async_container_client")?.unbind(), - }) - } - - fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult { - 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::()? - && self.container_name == sync_client.getattr("container_name")?.extract::()?) - } - - 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 { - Ok(Self { - reference: backend.getattr("s3_client")?.unbind(), - }) - } - - fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult { - 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 { - 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 { - 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 { - 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::()? != 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 { - 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) - } -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native/identity.rs b/litellm-rust/crates/python-bridge/src/cache/native/identity.rs deleted file mode 100644 index 3d4868dbe7b..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/native/identity.rs +++ /dev/null @@ -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, - default_ttl: Option, - }, - Redis { - topology: RedisTopology, - namespace: Option, - default_ttl: Option, - }, - S3 { - bucket: String, - key_prefix: String, - region: String, - endpoint: Option, - }, - Gcs { - bucket_name: String, - key_prefix: String, - path_service_account: Option, - }, - 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") - ); - } -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native/mod.rs b/litellm-rust/crates/python-bridge/src/cache/native/mod.rs index ad117dd9eaf..da1460f90ed 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native/mod.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/mod.rs @@ -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; diff --git a/litellm-rust/crates/python-bridge/src/cache/native/request.rs b/litellm-rust/crates/python-bridge/src/cache/native/request.rs deleted file mode 100644 index c009882a4ad..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/native/request.rs +++ /dev/null @@ -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, - ttl_seconds: Option, - max_age_seconds: Option, - messages: Option, - input: Option, - metadata: Option, - litellm_metadata: Option, - litellm_params: Option, - scope: Option, -} - -#[derive(Clone)] -pub(in crate::cache) struct NativeRequest { - pub(super) key: CacheKeyInput, - pub(super) controls: CacheControls, - pub(super) ttl: Option, - pub(super) max_age: Option, - pub(super) messages: Option, - pub(super) input: Option, - pub(super) metadata: Option, - pub(super) litellm_metadata: Option, - pub(super) litellm_params: Option, - pub(super) scope: Option, -} - -impl NativeRequest { - pub(super) fn exact(&self) -> ResponseCacheRequest { - 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 { - 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 { - self.semantic_with(semantic_key(self, scope), Some(scope.to_owned())) - } - - fn semantic_with( - &self, - key: CacheKeyInput, - scope: Option, - ) -> ResponseCacheRequest { - 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 { - let input: RequestInput = from_py(value)?; - request_input(input) -} - -fn request_input(input: RequestInput) -> PyResult { - let controls = input.controls.unwrap_or_else(|| { - ResponseCacheRequest::::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> { - from_py::>(value)? - .into_iter() - .map(request_input) - .collect() -} - -pub(super) fn duration(seconds: f64) -> PyResult { - 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") - ); - } -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native/semantic.rs b/litellm-rust/crates/python-bridge/src/cache/native/semantic.rs deleted file mode 100644 index 8d9bf270be0..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/native/semantic.rs +++ /dev/null @@ -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)>, - 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)> { - 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 { - 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>) -> PyResult { - 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::(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, Error>, - ) -> PyResult { - 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 { - 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), - Semantic(SemanticReply), -} - -impl ExecutionBody for SemanticExecution { - fn resume(&mut self, result: Option>>) -> PyResult { - Python::attach(|py| self.resume_py(py, result)) - } - - fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - self.embedder.traverse(visit) - } -} - -pub(super) fn drive(py: Python<'_>, body: SemanticExecution) -> PyResult> { - Execution::new(body, crate::lifecycle::binding).into_coroutine(py) -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native/v2.rs b/litellm-rust/crates/python-bridge/src/cache/native/v2.rs index 0dd70a042e9..82771644351 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native/v2.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/v2.rs @@ -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> { 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> { @@ -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() +} diff --git a/litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md b/litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md index ca4bd68dd38..81a0997f02a 100644 --- a/litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md +++ b/litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md @@ -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 diff --git a/litellm-rust/crates/python-bridge/src/cache/python/callback.rs b/litellm-rust/crates/python-bridge/src/cache/python/callback.rs deleted file mode 100644 index fd8ebe508bc..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/python/callback.rs +++ /dev/null @@ -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); - -impl PythonCallback { - pub(in crate::cache) fn new(object: Py) -> Self { - Self(object) - } - - pub(in crate::cache) fn lookup<'py>( - &self, - py: Python<'py>, - kwargs: Option<&Bound<'py, PyDict>>, - ) -> PyResult> { - 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> { - 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> { - 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> { - 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> { - let awaitables = batch_callback_kwargs(requests, kwargs)? - .iter() - .map(|kwargs| { - self.0 - .bind(py) - .call_method("async_get_cache", (), Some(kwargs)) - }) - .collect::>>()?; - 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> { - 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> { - 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> { - 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>> { - 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::()?)) - .collect::>>()?; - if kwargs.len() != requests.len()? { - return Err(PyValueError::new_err( - "batch cache requests and callback_kwargs must have equal lengths", - )); - } - Ok(kwargs) -} diff --git a/litellm-rust/crates/python-bridge/src/cache/python/mod.rs b/litellm-rust/crates/python-bridge/src/cache/python/mod.rs index 25550c2ac73..973dfe89f1f 100644 --- a/litellm-rust/crates/python-bridge/src/cache/python/mod.rs +++ b/litellm-rust/crates/python-bridge/src/cache/python/mod.rs @@ -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; diff --git a/litellm-rust/crates/python-bridge/src/cache/runtime.rs b/litellm-rust/crates/python-bridge/src/cache/runtime.rs deleted file mode 100644 index eec82f2ba4c..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/runtime.rs +++ /dev/null @@ -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, - 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> { - 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 { - 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 { - 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::>()?; - 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 { - 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> { - 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> { - 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> { - 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> { - 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> { - 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> { - 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> { - 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> { - self.check_process()?; - match &self.binding { - CacheBinding::Disabled => ready_none(py), - CacheBinding::Native(service) => { - let requests = self::requests(requests)?; - let responses: Vec = 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> { - 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> { - 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(()) - } -} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 6659be5160e..cf1080a3881 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -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::(), )?; - dict.set_item("_ResponseCacheRuntime", py.get_type::())?; dict.set_item( "_SecretManagerRuntime", py.get_type::(), diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 4070f41fd32..888adffda7f 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -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 diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 908321cdf96..4f9a0b2c492 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -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 diff --git a/litellm/rust_bridge/response_cache.py b/litellm/rust_bridge/response_cache.py deleted file mode 100644 index f100d9eff57..00000000000 --- a/litellm/rust_bridge/response_cache.py +++ /dev/null @@ -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)} diff --git a/tests/test_litellm_rust/cache/conftest.py b/tests/test_litellm_rust/cache/conftest.py index 07ccd0fde2b..35f69683986 100644 --- a/tests/test_litellm_rust/cache/conftest.py +++ b/tests/test_litellm_rust/cache/conftest.py @@ -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() diff --git a/tests/test_litellm_rust/cache/test_azure_blob.py b/tests/test_litellm_rust/cache/test_azure_blob.py deleted file mode 100644 index 064458ae9b0..00000000000 --- a/tests/test_litellm_rust/cache/test_azure_blob.py +++ /dev/null @@ -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() diff --git a/tests/test_litellm_rust/cache/test_disk.py b/tests/test_litellm_rust/cache/test_disk.py deleted file mode 100644 index e2faaa50221..00000000000 --- a/tests/test_litellm_rust/cache/test_disk.py +++ /dev/null @@ -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], - } diff --git a/tests/test_litellm_rust/cache/test_facade.py b/tests/test_litellm_rust/cache/test_facade.py deleted file mode 100644 index 1b91b7edb0c..00000000000 --- a/tests/test_litellm_rust/cache/test_facade.py +++ /dev/null @@ -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 diff --git a/tests/test_litellm_rust/cache/test_gcs.py b/tests/test_litellm_rust/cache/test_gcs.py deleted file mode 100644 index 5b81e48bc1d..00000000000 --- a/tests/test_litellm_rust/cache/test_gcs.py +++ /dev/null @@ -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)) diff --git a/tests/test_litellm_rust/cache/test_qdrant_semantic.py b/tests/test_litellm_rust/cache/test_qdrant_semantic.py deleted file mode 100644 index 529ec66bdd8..00000000000 --- a/tests/test_litellm_rust/cache/test_qdrant_semantic.py +++ /dev/null @@ -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"} diff --git a/tests/test_litellm_rust/cache/test_redis.py b/tests/test_litellm_rust/cache/test_redis.py deleted file mode 100644 index 881db9b1f2e..00000000000 --- a/tests/test_litellm_rust/cache/test_redis.py +++ /dev/null @@ -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 diff --git a/tests/test_litellm_rust/cache/test_redis_semantic.py b/tests/test_litellm_rust/cache/test_redis_semantic.py deleted file mode 100644 index a8b0c174d7e..00000000000 --- a/tests/test_litellm_rust/cache/test_redis_semantic.py +++ /dev/null @@ -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"} diff --git a/tests/test_litellm_rust/cache/test_rollout.py b/tests/test_litellm_rust/cache/test_rollout.py deleted file mode 100644 index a290857c8bb..00000000000 --- a/tests/test_litellm_rust/cache/test_rollout.py +++ /dev/null @@ -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 diff --git a/tests/test_litellm_rust/cache/test_s3.py b/tests/test_litellm_rust/cache/test_s3.py deleted file mode 100644 index d7b36ad1b2e..00000000000 --- a/tests/test_litellm_rust/cache/test_s3.py +++ /dev/null @@ -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) diff --git a/tests/test_litellm_rust/cache/test_valkey_semantic.py b/tests/test_litellm_rust/cache/test_valkey_semantic.py deleted file mode 100644 index ebd00695d95..00000000000 --- a/tests/test_litellm_rust/cache/test_valkey_semantic.py +++ /dev/null @@ -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"} diff --git a/tests/test_litellm_rust/support/cache.py b/tests/test_litellm_rust/support/cache.py index 0498eb1d10d..564578478cd 100644 --- a/tests/test_litellm_rust/support/cache.py +++ b/tests/test_litellm_rust/support/cache.py @@ -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): diff --git a/tests/test_litellm_rust/support/s3_stub.py b/tests/test_litellm_rust/support/s3_stub.py deleted file mode 100644 index 5a683fb78f3..00000000000 --- a/tests/test_litellm_rust/support/s3_stub.py +++ /dev/null @@ -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'NoSuchKey' - 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)