mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor(cache): remove dead Cache._native_cache runtime path (#44785)
The native response-cache runtime attached through Cache._native_cache is unreachable since the V2 cache replaced it. Drop the Python branches and wrapper, the _ResponseCacheRuntime pyclass and its backend/activation/ semantic modules, the python-bridge deps only they used, and the tests and fixtures dedicated to that path. Co-authored-by: Yujong Lee <yujong@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
a983b2a5e7
commit
5f9eff55f8
37 changed files with 28 additions and 7875 deletions
10
litellm-rust/Cargo.lock
generated
10
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,13 +0,0 @@
|
|||
use litellm_host_python::{ready_future, to_py};
|
||||
use pyo3::prelude::*;
|
||||
|
||||
pub(super) fn ready_none(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
|
||||
ready_value(py, &())
|
||||
}
|
||||
|
||||
pub(super) fn ready_value<'py, T: serde::Serialize>(
|
||||
py: Python<'py>,
|
||||
value: &T,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
ready_future(py, to_py(py, value)?.bind(py))
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,112 +0,0 @@
|
|||
use crate::cache::cache_error;
|
||||
use crate::execution::run_sync_value;
|
||||
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
|
||||
use litellm_cache_redis_semantic::RedisSemanticConfig;
|
||||
use litellm_host_python::release_gil;
|
||||
use litellm_http::ClientVariant;
|
||||
use pyo3::prelude::*;
|
||||
|
||||
use super::{
|
||||
backend::NativeResponseCache,
|
||||
config::{CacheBackendConfig, NativeCacheConfig, UnsupportedCacheConfig},
|
||||
embedder::PythonEmbedder,
|
||||
};
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
use crate::http::host_client;
|
||||
|
||||
fn declined(reason: UnsupportedCacheConfig) -> PyErr {
|
||||
RustBridgeDeclined::new_err(reason.message())
|
||||
}
|
||||
|
||||
/// Builds the native backend a `Cache` facade's projected configuration describes. `backend` is
|
||||
/// the facade's `.cache` object, which owns embedding for the Python-embedded semantic caches.
|
||||
pub(in crate::cache) fn activate(
|
||||
py: Python<'_>,
|
||||
backend: &Bound<'_, PyAny>,
|
||||
config: NativeCacheConfig,
|
||||
) -> PyResult<NativeResponseCache> {
|
||||
let policy = config.policy;
|
||||
let service = match config.backend {
|
||||
CacheBackendConfig::Memory(memory) => {
|
||||
NativeResponseCache::memory(memory.capacity, memory.default_ttl, memory.max_entry_bytes)
|
||||
}
|
||||
CacheBackendConfig::Redis(redis) => {
|
||||
let url = redis.connection.native_url().map_err(declined)?;
|
||||
let flush_size = policy.redis_flush_size.map(|_| redis.flush_size);
|
||||
release_gil(py, move || {
|
||||
NativeResponseCache::redis(
|
||||
&url,
|
||||
&redis.topology,
|
||||
Some(redis.default_ttl),
|
||||
redis.namespace,
|
||||
)
|
||||
})
|
||||
.map_err(cache_error)?
|
||||
.with_redis_flush_size(flush_size)
|
||||
}
|
||||
CacheBackendConfig::S3(s3) => {
|
||||
let http = host_client(py, ClientVariant::NoRedirect)?;
|
||||
run_sync_value(
|
||||
py,
|
||||
async move { Ok(NativeResponseCache::s3(*s3, http).await) },
|
||||
)?
|
||||
}
|
||||
CacheBackendConfig::Gcs(gcs) => NativeResponseCache::gcs(
|
||||
GcsConfig {
|
||||
bucket_name: gcs.bucket_name,
|
||||
gcs_path: Some(gcs.key_prefix),
|
||||
path_service_account: gcs.path_service_account,
|
||||
endpoint: DEFAULT_ENDPOINT.to_owned(),
|
||||
},
|
||||
host_client(py, ClientVariant::NoRedirect)?,
|
||||
None,
|
||||
),
|
||||
CacheBackendConfig::Disk(disk) => {
|
||||
release_gil(py, move || NativeResponseCache::disk(&disk.directory))
|
||||
.map_err(cache_error)?
|
||||
}
|
||||
CacheBackendConfig::AzureBlob(azure) => {
|
||||
let http = host_client(py, ClientVariant::NoRedirect)?;
|
||||
run_sync_value(py, async move {
|
||||
NativeResponseCache::azure_blob(&azure.account_url, &azure.container, http)
|
||||
.await
|
||||
.map_err(cache_error)
|
||||
})?
|
||||
}
|
||||
CacheBackendConfig::RedisSemantic(semantic) => {
|
||||
let url = semantic.native_url().map_err(declined)?.to_owned();
|
||||
let embedder = PythonEmbedder::new(backend.clone().unbind());
|
||||
let semantic_config = RedisSemanticConfig {
|
||||
index_name: semantic.index_name,
|
||||
similarity_threshold: semantic.similarity_threshold as f32,
|
||||
};
|
||||
release_gil(py, move || {
|
||||
NativeResponseCache::redis_semantic(&url, embedder, semantic_config)
|
||||
})
|
||||
.map_err(cache_error)?
|
||||
}
|
||||
CacheBackendConfig::ValkeySemantic(valkey) => {
|
||||
let url = valkey.connection.native_url().map_err(declined)?;
|
||||
let embedder = PythonEmbedder::new(backend.clone().unbind());
|
||||
release_gil(py, move || {
|
||||
NativeResponseCache::valkey_semantic(
|
||||
&url,
|
||||
valkey.similarity_threshold,
|
||||
valkey.index_name,
|
||||
embedder,
|
||||
)
|
||||
})
|
||||
.map_err(cache_error)?
|
||||
}
|
||||
CacheBackendConfig::QdrantSemantic(qdrant) => {
|
||||
let client = host_client(py, ClientVariant::Provider)?;
|
||||
run_sync_value(py, async move {
|
||||
let runtime = tokio::runtime::Handle::current();
|
||||
NativeResponseCache::qdrant_semantic(*qdrant, client, runtime)
|
||||
.await
|
||||
.map_err(cache_error)
|
||||
})?
|
||||
}
|
||||
};
|
||||
Ok(service.with_scope(policy.semantic_cache_scope))
|
||||
}
|
||||
|
|
@ -1,690 +0,0 @@
|
|||
use crate::cache::cache_error;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use litellm_cache::{CacheCodec, CacheConnectionResult, Error, semantic::SemanticLookup};
|
||||
use litellm_cache_azure_blob::AzureBlobCache;
|
||||
use litellm_cache_disk::DiskCache;
|
||||
use litellm_cache_gcs::{GcsCache, GcsConfig, StaticTokenSource};
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_qdrant_semantic::{OpenAiEmbedder, QdrantSemanticCache};
|
||||
use litellm_cache_redis::{RedisCache, RedisTopology};
|
||||
use litellm_cache_redis_semantic::{RedisSemanticCache, RedisSemanticConfig};
|
||||
use litellm_cache_response::{
|
||||
ConnectionProbe, ExactResponseCache, PartialHits, ResponseCache, ResponseCacheCodec,
|
||||
WriteBuffer,
|
||||
};
|
||||
use litellm_cache_s3::{S3Cache, S3CacheConfig};
|
||||
use litellm_cache_valkey_semantic::{ValkeySemanticCache, ValkeySemanticConfig};
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
config::QdrantSemanticCacheConfig,
|
||||
embedder::PythonEmbedder,
|
||||
identity::BackendIdentity,
|
||||
request::{NativeRequest, now},
|
||||
semantic::{EmbeddingFailure, SemanticExecution, SemanticOperation, drive},
|
||||
};
|
||||
|
||||
/// What the Python embedder receives for one semantic request.
|
||||
pub(in crate::cache) struct EmbeddingInput {
|
||||
pub(in crate::cache) prompt: String,
|
||||
pub(in crate::cache) metadata: Option<Value>,
|
||||
}
|
||||
|
||||
/// An exact-match backend behind one pointer, with the identity its facade must reproduce.
|
||||
pub(in crate::cache) struct ExactService {
|
||||
cache: Arc<dyn ExactResponseCache>,
|
||||
probe: Option<Arc<dyn ConnectionProbe>>,
|
||||
buffer: Option<WriteBuffer>,
|
||||
identity: BackendIdentity,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(in crate::cache) enum NativeResponseCache {
|
||||
Exact(Arc<ExactService>),
|
||||
ValkeySemantic {
|
||||
cache: Arc<ResponseCache<ValkeySemanticCache<PythonEmbedder, ResponseCacheCodec>>>,
|
||||
embedder: PythonEmbedder,
|
||||
scope: String,
|
||||
},
|
||||
RedisSemantic {
|
||||
cache: Arc<ResponseCache<RedisSemanticCache<PythonEmbedder, ResponseCacheCodec>>>,
|
||||
embedder: PythonEmbedder,
|
||||
},
|
||||
QdrantSemantic(Arc<ResponseCache<QdrantSemanticCache<OpenAiEmbedder, ResponseCacheCodec>>>),
|
||||
}
|
||||
|
||||
impl NativeResponseCache {
|
||||
pub fn memory(capacity: usize, ttl: Duration, max_entry_bytes: usize) -> Self {
|
||||
let backend = InMemoryCache::with_clock_and_size_measurement(
|
||||
Some(capacity),
|
||||
Some(ttl),
|
||||
Some(max_entry_bytes),
|
||||
Some(Arc::new(|entry| {
|
||||
ResponseCacheCodec.encode(entry).map(|bytes| bytes.len())
|
||||
})),
|
||||
now,
|
||||
);
|
||||
let identity = BackendIdentity::Memory {
|
||||
capacity: backend.max_size_in_memory(),
|
||||
max_entry_bytes: backend.max_entry_bytes(),
|
||||
default_ttl: None,
|
||||
};
|
||||
Self::exact(ResponseCache::new(Arc::new(backend)), identity)
|
||||
}
|
||||
|
||||
pub fn redis(
|
||||
url: &str,
|
||||
topology: &RedisTopology,
|
||||
ttl: Option<Duration>,
|
||||
namespace: Option<String>,
|
||||
) -> Result<Self, Error> {
|
||||
let backend =
|
||||
RedisCache::connect(url, topology, ttl, ResponseCacheCodec)?.with_namespace(namespace);
|
||||
let identity = BackendIdentity::Redis {
|
||||
topology: backend.topology().clone(),
|
||||
namespace: backend.namespace().map(str::to_owned),
|
||||
default_ttl: None,
|
||||
};
|
||||
Ok(Self::exact_probed(
|
||||
ResponseCache::new(Arc::new(backend)),
|
||||
identity,
|
||||
))
|
||||
}
|
||||
|
||||
pub async fn s3(config: S3CacheConfig, http: litellm_http::Client) -> Self {
|
||||
let runtime = tokio::runtime::Handle::current();
|
||||
let backend = S3Cache::new(config, http, ResponseCacheCodec, runtime);
|
||||
let identity = BackendIdentity::S3 {
|
||||
bucket: backend.bucket().to_owned(),
|
||||
key_prefix: backend.key_prefix().to_owned(),
|
||||
region: backend.region().to_owned(),
|
||||
endpoint: backend.endpoint().map(str::to_owned),
|
||||
};
|
||||
Self::exact(ResponseCache::new(Arc::new(backend)), identity)
|
||||
}
|
||||
|
||||
pub fn disk(directory: impl AsRef<std::path::Path>) -> Result<Self, Error> {
|
||||
let backend = DiskCache::open(directory, ResponseCacheCodec)?;
|
||||
let identity = BackendIdentity::Disk {
|
||||
directory: backend.directory().to_path_buf(),
|
||||
};
|
||||
Ok(Self::exact(ResponseCache::new(Arc::new(backend)), identity))
|
||||
}
|
||||
|
||||
pub fn gcs(config: GcsConfig, client: litellm_http::Client, token: Option<String>) -> Self {
|
||||
let backend = match token {
|
||||
Some(token) => GcsCache::with_token_source(
|
||||
config,
|
||||
client,
|
||||
ResponseCacheCodec,
|
||||
Arc::new(StaticTokenSource(token)),
|
||||
),
|
||||
None => GcsCache::new(config, client, ResponseCacheCodec),
|
||||
};
|
||||
let identity = BackendIdentity::Gcs {
|
||||
bucket_name: backend.bucket_name().to_owned(),
|
||||
key_prefix: backend.key_prefix().to_owned(),
|
||||
path_service_account: backend.path_service_account().map(str::to_owned),
|
||||
};
|
||||
Self::exact(ResponseCache::new(Arc::new(backend)), identity)
|
||||
}
|
||||
|
||||
pub async fn azure_blob(
|
||||
account_url: &str,
|
||||
container: &str,
|
||||
http: litellm_http::Client,
|
||||
) -> Result<Self, Error> {
|
||||
let backend = AzureBlobCache::connect(
|
||||
account_url,
|
||||
container,
|
||||
http,
|
||||
ResponseCacheCodec,
|
||||
tokio::runtime::Handle::current(),
|
||||
)
|
||||
.await?;
|
||||
let identity = BackendIdentity::AzureBlob {
|
||||
account_url: backend.account_url().to_owned(),
|
||||
container: backend.container_name().to_owned(),
|
||||
};
|
||||
Ok(Self::exact(ResponseCache::new(Arc::new(backend)), identity))
|
||||
}
|
||||
|
||||
/// Wraps a built exact backend; the TTL a facade must match comes from the built cache.
|
||||
fn exact<B>(cache: ResponseCache<B>, identity: BackendIdentity) -> Self
|
||||
where
|
||||
ResponseCache<B>: ExactResponseCache + 'static,
|
||||
B: litellm_cache::BaseCache<Value = litellm_cache_response::CacheEntry>,
|
||||
B::Context: Default + PartialEq,
|
||||
{
|
||||
Self::exact_service(Arc::new(cache), None, identity)
|
||||
}
|
||||
|
||||
/// Wraps an exact backend whose Python class defines `test_connection`.
|
||||
fn exact_probed<B>(cache: ResponseCache<B>, identity: BackendIdentity) -> Self
|
||||
where
|
||||
ResponseCache<B>: ExactResponseCache + ConnectionProbe + 'static,
|
||||
B: litellm_cache::BaseCache<Value = litellm_cache_response::CacheEntry>,
|
||||
B::Context: Default + PartialEq,
|
||||
{
|
||||
let cache = Arc::new(cache);
|
||||
Self::exact_service(cache.clone(), Some(cache), identity)
|
||||
}
|
||||
|
||||
fn exact_service(
|
||||
cache: Arc<dyn ExactResponseCache>,
|
||||
probe: Option<Arc<dyn ConnectionProbe>>,
|
||||
identity: BackendIdentity,
|
||||
) -> Self {
|
||||
let default_ttl = cache.default_ttl();
|
||||
let identity = match identity {
|
||||
BackendIdentity::Memory {
|
||||
capacity,
|
||||
max_entry_bytes,
|
||||
..
|
||||
} => BackendIdentity::Memory {
|
||||
capacity,
|
||||
max_entry_bytes,
|
||||
default_ttl,
|
||||
},
|
||||
BackendIdentity::Redis {
|
||||
topology,
|
||||
namespace,
|
||||
..
|
||||
} => BackendIdentity::Redis {
|
||||
topology,
|
||||
namespace,
|
||||
default_ttl,
|
||||
},
|
||||
other => other,
|
||||
};
|
||||
Self::Exact(Arc::new(ExactService {
|
||||
cache,
|
||||
probe,
|
||||
buffer: None,
|
||||
identity,
|
||||
}))
|
||||
}
|
||||
|
||||
pub fn valkey_semantic(
|
||||
url: &str,
|
||||
similarity_threshold: f64,
|
||||
index_name: String,
|
||||
embedder: PythonEmbedder,
|
||||
) -> Result<Self, Error> {
|
||||
let backend = ValkeySemanticCache::new(
|
||||
url,
|
||||
embedder.clone(),
|
||||
ResponseCacheCodec,
|
||||
ValkeySemanticConfig {
|
||||
similarity_threshold,
|
||||
index_name,
|
||||
},
|
||||
)?;
|
||||
Ok(Self::ValkeySemantic {
|
||||
cache: Arc::new(ResponseCache::new(Arc::new(backend))),
|
||||
embedder,
|
||||
scope: String::from("key"),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn redis_semantic(
|
||||
url: &str,
|
||||
embedder: PythonEmbedder,
|
||||
config: RedisSemanticConfig,
|
||||
) -> Result<Self, Error> {
|
||||
let backend = RedisSemanticCache::new(url, embedder.clone(), ResponseCacheCodec, config)?;
|
||||
Ok(Self::RedisSemantic {
|
||||
cache: Arc::new(ResponseCache::new(Arc::new(backend))),
|
||||
embedder,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn qdrant_semantic(
|
||||
config: QdrantSemanticCacheConfig,
|
||||
client: litellm_http::Client,
|
||||
runtime: tokio::runtime::Handle,
|
||||
) -> Result<Self, Error> {
|
||||
let qdrant = qdrant_client::Qdrant::from_url(&config.grpc_url)
|
||||
.skip_compatibility_check()
|
||||
.api_key(config.api_key.as_deref())
|
||||
.build()
|
||||
.map_err(|_| Error::Unavailable)?;
|
||||
let qdrant_config = config.to_qdrant_config();
|
||||
let embedder = OpenAiEmbedder::new(client, config.embedding);
|
||||
let cache = QdrantSemanticCache::connect(
|
||||
qdrant,
|
||||
embedder,
|
||||
ResponseCacheCodec,
|
||||
qdrant_config,
|
||||
runtime,
|
||||
)
|
||||
.await?;
|
||||
Ok(Self::QdrantSemantic(Arc::new(ResponseCache::new(
|
||||
Arc::new(cache),
|
||||
))))
|
||||
}
|
||||
|
||||
pub fn identity(&self) -> BackendIdentity {
|
||||
match self {
|
||||
Self::Exact(service) => service.identity.clone(),
|
||||
Self::ValkeySemantic { cache, .. } => BackendIdentity::ValkeySemantic {
|
||||
index_name: cache.backend().index_name().to_owned(),
|
||||
similarity_threshold: cache.backend().similarity_threshold(),
|
||||
},
|
||||
Self::RedisSemantic { cache, .. } => BackendIdentity::RedisSemantic {
|
||||
index_name: cache.backend().index_name().to_owned(),
|
||||
similarity_threshold: cache.backend().similarity_threshold(),
|
||||
},
|
||||
Self::QdrantSemantic(cache) => BackendIdentity::QdrantSemantic {
|
||||
collection_name: cache.backend().collection_name().to_owned(),
|
||||
similarity_threshold: cache.backend().similarity_threshold(),
|
||||
vector_size: cache.backend().vector_size(),
|
||||
embedding_model: cache.backend().embedder().model().to_owned(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_redis_flush_size(self, flush_size: Option<usize>) -> Self {
|
||||
match self {
|
||||
Self::Exact(service) if matches!(service.identity, BackendIdentity::Redis { .. }) => {
|
||||
Self::Exact(Arc::new(ExactService {
|
||||
cache: Arc::clone(&service.cache),
|
||||
probe: service.probe.clone(),
|
||||
buffer: flush_size.map(WriteBuffer::new),
|
||||
identity: service.identity.clone(),
|
||||
}))
|
||||
}
|
||||
value => value,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_scope(self, scope: String) -> Self {
|
||||
match self {
|
||||
Self::ValkeySemantic {
|
||||
cache, embedder, ..
|
||||
} => Self::ValkeySemantic {
|
||||
cache,
|
||||
embedder,
|
||||
scope,
|
||||
},
|
||||
value => value,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn embedder_object(&self) -> Option<&Py<PyAny>> {
|
||||
match self {
|
||||
Self::RedisSemantic { embedder, .. } => Some(embedder.object()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// The prompt and metadata this backend would embed for `request`, if it has a prompt.
|
||||
pub(in crate::cache) fn embedding_input(
|
||||
&self,
|
||||
request: &NativeRequest,
|
||||
) -> Option<EmbeddingInput> {
|
||||
let context = match self {
|
||||
Self::ValkeySemantic { scope, .. } => request.scoped_semantic(scope).context,
|
||||
Self::RedisSemantic { .. } => request.semantic().context,
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => return None,
|
||||
};
|
||||
let prompt = litellm_cache::semantic::prompt_from_context(&context)?;
|
||||
Some(EmbeddingInput {
|
||||
prompt,
|
||||
metadata: context.metadata,
|
||||
})
|
||||
}
|
||||
|
||||
/// Drives a semantic operation whose embedding comes from Python.
|
||||
fn python_semantic<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
operation: SemanticOperation,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let (embedder, failure) = match self {
|
||||
Self::ValkeySemantic { embedder, .. } => (embedder, EmbeddingFailure::Propagate),
|
||||
Self::RedisSemantic { embedder, .. } => (embedder, EmbeddingFailure::Unavailable),
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
return Err(pyo3::exceptions::PyRuntimeError::new_err(
|
||||
"semantic execution requires a Python-embedded backend",
|
||||
));
|
||||
}
|
||||
};
|
||||
drive(
|
||||
py,
|
||||
SemanticExecution::new(self.clone(), embedder.clone(), failure, operation),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn lookup(&self, request: &NativeRequest, now: Duration) -> Result<Option<Value>, Error> {
|
||||
match self {
|
||||
Self::Exact(service) => service.cache.lookup(&request.exact(), now),
|
||||
Self::ValkeySemantic { cache, scope, .. } => {
|
||||
cache.lookup(&request.scoped_semantic(scope), now)
|
||||
}
|
||||
Self::RedisSemantic { cache, .. } => cache.lookup(&request.semantic(), now),
|
||||
Self::QdrantSemantic(cache) => cache.lookup(&request.semantic(), now),
|
||||
}
|
||||
}
|
||||
|
||||
/// `lookup` plus the similarity Python's semantic backend writes to the request metadata.
|
||||
/// Exact backends report none.
|
||||
pub fn lookup_semantic(
|
||||
&self,
|
||||
request: &NativeRequest,
|
||||
now: Duration,
|
||||
) -> Result<SemanticLookup<Value>, Error> {
|
||||
match self {
|
||||
Self::Exact(service) => service
|
||||
.cache
|
||||
.lookup(&request.exact(), now)
|
||||
.map(exact_lookup),
|
||||
Self::ValkeySemantic { cache, scope, .. } => {
|
||||
redis_family(cache.lookup_semantic(&request.scoped_semantic(scope), now))
|
||||
}
|
||||
Self::RedisSemantic { cache, .. } => {
|
||||
redis_family(cache.lookup_semantic(&request.semantic(), now))
|
||||
}
|
||||
Self::QdrantSemantic(cache) => cache.lookup_semantic(&request.semantic(), now),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn store(
|
||||
&self,
|
||||
request: &NativeRequest,
|
||||
response: Value,
|
||||
now: Duration,
|
||||
) -> Result<(), Error> {
|
||||
match self {
|
||||
Self::Exact(service) => service.cache.store(&request.exact(), response, now),
|
||||
Self::ValkeySemantic { cache, scope, .. } => {
|
||||
cache.store(&request.scoped_semantic(scope), response, now)
|
||||
}
|
||||
Self::RedisSemantic { cache, .. } => cache.store(&request.semantic(), response, now),
|
||||
Self::QdrantSemantic(cache) => cache.store(&request.semantic(), response, now),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn lookup_batch(
|
||||
&self,
|
||||
requests: &[NativeRequest],
|
||||
now: Duration,
|
||||
) -> Result<PartialHits, Error> {
|
||||
match self {
|
||||
Self::Exact(service) => service.cache.lookup_batch(&exact_requests(requests), now),
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => {
|
||||
Err(Error::UnsupportedOperation)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn async_lookup(
|
||||
&self,
|
||||
request: &NativeRequest,
|
||||
now: Duration,
|
||||
) -> Result<Option<Value>, Error> {
|
||||
match self {
|
||||
Self::Exact(service) => service.cache.async_lookup(&request.exact(), now).await,
|
||||
Self::ValkeySemantic { cache, scope, .. } => {
|
||||
cache
|
||||
.async_lookup(&request.scoped_semantic(scope), now)
|
||||
.await
|
||||
}
|
||||
Self::RedisSemantic { cache, .. } => cache.async_lookup(&request.semantic(), now).await,
|
||||
Self::QdrantSemantic(cache) => cache.async_lookup(&request.semantic(), now).await,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn async_lookup_semantic(
|
||||
&self,
|
||||
request: &NativeRequest,
|
||||
now: Duration,
|
||||
) -> Result<SemanticLookup<Value>, Error> {
|
||||
match self {
|
||||
Self::Exact(service) => service
|
||||
.cache
|
||||
.async_lookup(&request.exact(), now)
|
||||
.await
|
||||
.map(exact_lookup),
|
||||
Self::ValkeySemantic { cache, scope, .. } => redis_family(
|
||||
cache
|
||||
.async_lookup_semantic(&request.scoped_semantic(scope), now)
|
||||
.await,
|
||||
),
|
||||
Self::RedisSemantic { cache, .. } => {
|
||||
redis_family(cache.async_lookup_semantic(&request.semantic(), now).await)
|
||||
}
|
||||
Self::QdrantSemantic(cache) => {
|
||||
cache.async_lookup_semantic(&request.semantic(), now).await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn async_lookup_semantic_py<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
request: NativeRequest,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
service
|
||||
.async_lookup_semantic(&request, now())
|
||||
.await
|
||||
.map(SemanticReply::from)
|
||||
},
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
|
||||
self.python_semantic(py, SemanticOperation::LookupSemantic(request))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn async_lookup_py<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
request: NativeRequest,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { service.async_lookup(&request, now()).await },
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
|
||||
self.python_semantic(py, SemanticOperation::Lookup(request))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn async_store(
|
||||
&self,
|
||||
request: &NativeRequest,
|
||||
response: Value,
|
||||
now: Duration,
|
||||
) -> Result<(), Error> {
|
||||
match self {
|
||||
Self::Exact(service) => match &service.buffer {
|
||||
None => {
|
||||
service
|
||||
.cache
|
||||
.async_store(&request.exact(), response, now)
|
||||
.await
|
||||
}
|
||||
Some(buffer) => {
|
||||
buffer
|
||||
.async_store(service.cache.as_ref(), &request.exact(), response, now)
|
||||
.await
|
||||
}
|
||||
},
|
||||
Self::ValkeySemantic { cache, scope, .. } => {
|
||||
cache
|
||||
.async_store(&request.scoped_semantic(scope), response, now)
|
||||
.await
|
||||
}
|
||||
Self::RedisSemantic { cache, .. } => {
|
||||
cache.async_store(&request.semantic(), response, now).await
|
||||
}
|
||||
Self::QdrantSemantic(cache) => {
|
||||
cache.async_store(&request.semantic(), response, now).await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn async_store_py<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
request: NativeRequest,
|
||||
response: Value,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { service.async_store(&request, response, now()).await },
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
|
||||
self.python_semantic(py, SemanticOperation::Store(request, response))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn async_lookup_batch(
|
||||
&self,
|
||||
requests: &[NativeRequest],
|
||||
now: Duration,
|
||||
) -> Result<PartialHits, Error> {
|
||||
match self {
|
||||
Self::Exact(service) => {
|
||||
service
|
||||
.cache
|
||||
.async_lookup_batch(&exact_requests(requests), now)
|
||||
.await
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => {
|
||||
Err(Error::UnsupportedOperation)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn async_store_batch(
|
||||
&self,
|
||||
entries: Vec<(NativeRequest, Value)>,
|
||||
now: Duration,
|
||||
) -> Result<(), Error> {
|
||||
match self {
|
||||
Self::Exact(service) => {
|
||||
let entries = entries
|
||||
.into_iter()
|
||||
.map(|(request, value)| (request.exact(), value))
|
||||
.collect();
|
||||
service.cache.async_store_batch(entries, now).await
|
||||
}
|
||||
Self::ValkeySemantic { cache, scope, .. } => {
|
||||
let entries = entries
|
||||
.into_iter()
|
||||
.map(|(request, value)| (request.scoped_semantic(scope), value))
|
||||
.collect();
|
||||
cache.async_store_batch(entries, now).await
|
||||
}
|
||||
Self::RedisSemantic { .. } => Err(Error::UnsupportedOperation),
|
||||
Self::QdrantSemantic(cache) => {
|
||||
let entries = entries
|
||||
.into_iter()
|
||||
.map(|(request, value)| (request.semantic(), value))
|
||||
.collect();
|
||||
cache.async_store_batch(entries, now).await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn async_store_batch_py<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
entries: Vec<(NativeRequest, Value)>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { service.async_store_batch(entries, now()).await },
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
|
||||
self.python_semantic(py, SemanticOperation::StoreBatch(entries.into()))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn async_flush(&self) -> Result<(), Error> {
|
||||
match self {
|
||||
Self::Exact(service) => {
|
||||
if let Some(buffer) = &service.buffer {
|
||||
buffer.clear()?;
|
||||
}
|
||||
service.cache.async_flush().await
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => {
|
||||
Err(Error::UnsupportedOperation)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
|
||||
match self {
|
||||
Self::Exact(service) => match &service.probe {
|
||||
Some(probe) => probe.test_connection().await,
|
||||
None => Err(Error::UnsupportedOperation),
|
||||
},
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => {
|
||||
Err(Error::UnsupportedOperation)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn exact_requests(requests: &[NativeRequest]) -> Vec<litellm_cache_response::ResponseCacheRequest> {
|
||||
requests.iter().map(NativeRequest::exact).collect()
|
||||
}
|
||||
|
||||
/// What `lookup_semantic` hands Python: the response and the similarity to stamp, if any.
|
||||
#[derive(serde::Serialize)]
|
||||
pub(in crate::cache) struct SemanticReply(
|
||||
pub(in crate::cache) Option<Value>,
|
||||
pub(in crate::cache) Option<f64>,
|
||||
);
|
||||
|
||||
impl From<SemanticLookup<Value>> for SemanticReply {
|
||||
fn from(lookup: SemanticLookup<Value>) -> Self {
|
||||
Self(lookup.value, lookup.similarity)
|
||||
}
|
||||
}
|
||||
|
||||
fn exact_lookup(value: Option<Value>) -> SemanticLookup<Value> {
|
||||
SemanticLookup {
|
||||
value,
|
||||
similarity: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Python's Redis and Valkey semantic caches catch every lookup failure and stamp `0.0`.
|
||||
fn redis_family(
|
||||
lookup: Result<SemanticLookup<Value>, Error>,
|
||||
) -> Result<SemanticLookup<Value>, Error> {
|
||||
Ok(lookup.unwrap_or_else(|_| SemanticLookup::miss(Some(0.0))))
|
||||
}
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,129 +0,0 @@
|
|||
use std::future::Future;
|
||||
|
||||
use litellm_cache::Error;
|
||||
use litellm_host_python::to_py;
|
||||
use pyo3::{PyTraverseError, PyVisit, prelude::*, types::PyDict};
|
||||
use serde_json::Value;
|
||||
|
||||
tokio::task_local! {
|
||||
static PREPARED_EMBEDDING: Result<Vec<f32>, Error>;
|
||||
}
|
||||
|
||||
/// Runs `future` with the vector the Python embedder already produced, so the backend's
|
||||
/// `async_embed` never has to call back into Python from the runtime.
|
||||
pub(in crate::cache) fn with_prepared_embedding<F: Future>(
|
||||
vector: Result<Vec<f32>, Error>,
|
||||
future: F,
|
||||
) -> impl Future<Output = F::Output> {
|
||||
PREPARED_EMBEDDING.scope(vector, future)
|
||||
}
|
||||
|
||||
/// The Python object that owns embedding for a semantic backend.
|
||||
pub(in crate::cache) struct PythonEmbedder(Py<PyAny>);
|
||||
|
||||
impl Clone for PythonEmbedder {
|
||||
fn clone(&self) -> Self {
|
||||
Python::attach(|py| Self(self.0.clone_ref(py)))
|
||||
}
|
||||
}
|
||||
|
||||
impl PythonEmbedder {
|
||||
pub(in crate::cache) fn new(object: Py<PyAny>) -> Self {
|
||||
Self(object)
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn object(&self) -> &Py<PyAny> {
|
||||
&self.0
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.0)
|
||||
}
|
||||
|
||||
fn metadata_kwargs<'py>(
|
||||
py: Python<'py>,
|
||||
metadata: Option<&Value>,
|
||||
) -> PyResult<Bound<'py, PyDict>> {
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs.set_item("metadata", to_py(py, &metadata)?)?;
|
||||
Ok(kwargs)
|
||||
}
|
||||
|
||||
/// The awaitable of `_get_async_embedding(prompt, metadata=...)`, to run in the caller's loop.
|
||||
pub(in crate::cache) fn async_embedding(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
prompt: &str,
|
||||
metadata: Option<&Value>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let kwargs = Self::metadata_kwargs(py, metadata)?;
|
||||
self.0
|
||||
.bind(py)
|
||||
.call_method("_get_async_embedding", (prompt,), Some(&kwargs))
|
||||
.map(Bound::unbind)
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn extract(vector: Bound<'_, PyAny>) -> PyResult<Vec<f32>> {
|
||||
Ok(vector
|
||||
.extract::<Vec<f64>>()?
|
||||
.into_iter()
|
||||
.map(|value| value as f32)
|
||||
.collect())
|
||||
}
|
||||
|
||||
fn embed_sync(&self, prompt: &str, metadata: Option<&Value>) -> Result<Vec<f32>, Error> {
|
||||
Python::attach(|py| {
|
||||
let kwargs = Self::metadata_kwargs(py, metadata)?;
|
||||
Self::extract(self.0.bind(py).call_method(
|
||||
"_get_embedding",
|
||||
(prompt,),
|
||||
Some(&kwargs),
|
||||
)?)
|
||||
})
|
||||
.map_err(|_| Error::Unavailable)
|
||||
}
|
||||
|
||||
fn seeded_embedding() -> Result<Vec<f32>, Error> {
|
||||
PREPARED_EMBEDDING
|
||||
.try_with(Clone::clone)
|
||||
.unwrap_or(Err(Error::Unavailable))
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_cache::semantic::Embedder for PythonEmbedder {
|
||||
fn embed(&self, prompt: &str, metadata: Option<&Value>) -> Result<Vec<f32>, Error> {
|
||||
self.embed_sync(prompt, metadata)
|
||||
}
|
||||
|
||||
fn async_embed(
|
||||
&self,
|
||||
_prompt: &str,
|
||||
_metadata: Option<&Value>,
|
||||
) -> impl Future<Output = Result<Vec<f32>, Error>> + Send {
|
||||
std::future::ready(Self::seeded_embedding())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn async_embed_returns_the_seeded_vector_or_unavailable() {
|
||||
Python::initialize();
|
||||
let object = Python::attach(|py| py.None());
|
||||
let embedder = PythonEmbedder::new(object);
|
||||
let scoped_embedder = embedder.clone();
|
||||
let scoped = with_prepared_embedding(Ok(vec![0.25]), async move {
|
||||
litellm_cache::semantic::Embedder::async_embed(&scoped_embedder, "prompt", None).await
|
||||
});
|
||||
assert_eq!(scoped.await, Ok(vec![0.25]));
|
||||
let unscoped =
|
||||
litellm_cache::semantic::Embedder::async_embed(&embedder, "prompt", None).await;
|
||||
assert_eq!(unscoped, Err(Error::Unavailable));
|
||||
let valkey = with_prepared_embedding(Ok(vec![0.5]), async move {
|
||||
litellm_cache::semantic::Embedder::async_embed(&embedder, "prompt", None).await
|
||||
});
|
||||
assert_eq!(valkey.await, Ok(vec![0.5]));
|
||||
}
|
||||
}
|
||||
|
|
@ -1,502 +0,0 @@
|
|||
use litellm_cache_redis::RedisTopology;
|
||||
use litellm_host_python::from_py;
|
||||
use pyo3::{
|
||||
PyTraverseError, PyVisit,
|
||||
exceptions::PyTypeError,
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple, PyType},
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
backend::NativeResponseCache,
|
||||
config::{CacheConfigProjection, NativeCacheConfig},
|
||||
identity::BackendIdentity,
|
||||
};
|
||||
|
||||
struct ClassGuard {
|
||||
class: Py<PyType>,
|
||||
attributes: Vec<(String, Py<PyAny>)>,
|
||||
}
|
||||
|
||||
struct ObjectGuard {
|
||||
reference: Py<PyAny>,
|
||||
classes: Vec<ClassGuard>,
|
||||
config_names: &'static [&'static str],
|
||||
config: Vec<Value>,
|
||||
}
|
||||
|
||||
struct RedisPoolGuard {
|
||||
reference: Py<PyAny>,
|
||||
connection_class: Py<PyAny>,
|
||||
connection_kwargs: Py<PyAny>,
|
||||
max_connections: Option<usize>,
|
||||
client_name: &'static str,
|
||||
attributes: RedisPoolAttributes,
|
||||
}
|
||||
|
||||
struct DiskStoreGuard {
|
||||
reference: Py<PyAny>,
|
||||
directory: String,
|
||||
}
|
||||
|
||||
struct AzureBlobClientGuard {
|
||||
sync_client: Py<PyAny>,
|
||||
async_client: Py<PyAny>,
|
||||
url: String,
|
||||
container_name: String,
|
||||
}
|
||||
|
||||
struct S3ClientGuard {
|
||||
reference: Py<PyAny>,
|
||||
}
|
||||
|
||||
enum ConnectionGuard {
|
||||
None,
|
||||
RedisPool(RedisPoolGuard),
|
||||
AzureBlob(AzureBlobClientGuard),
|
||||
S3(S3ClientGuard),
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
struct RedisPoolAttributes {
|
||||
pool: &'static str,
|
||||
connection_class: &'static str,
|
||||
max_connections: Option<&'static str>,
|
||||
}
|
||||
|
||||
const STANDALONE_POOL: RedisPoolAttributes = RedisPoolAttributes {
|
||||
pool: "connection_pool",
|
||||
connection_class: "connection_class",
|
||||
max_connections: Some("max_connections"),
|
||||
};
|
||||
|
||||
const CLUSTER_POOL: RedisPoolAttributes = RedisPoolAttributes {
|
||||
pool: "nodes_manager",
|
||||
connection_class: "connection_pool_class",
|
||||
max_connections: None,
|
||||
};
|
||||
|
||||
const VALKEY_POOL: RedisPoolAttributes = STANDALONE_POOL;
|
||||
|
||||
/// Class-level defaults an instance overwrites with its own state rather than behavior:
|
||||
/// `Cache._native_cache` holds the runtime `Cache.__init__` resolved.
|
||||
const INSTANCE_STATE: &[&str] = &["_native_cache"];
|
||||
|
||||
pub(in crate::cache) struct FacadeGuard {
|
||||
outer: ObjectGuard,
|
||||
backend: ObjectGuard,
|
||||
disk_store: Option<DiskStoreGuard>,
|
||||
connection: ConnectionGuard,
|
||||
}
|
||||
|
||||
impl ObjectGuard {
|
||||
fn capture(
|
||||
py: Python<'_>,
|
||||
object: &Bound<'_, PyAny>,
|
||||
config_names: &'static [&'static str],
|
||||
) -> PyResult<Self> {
|
||||
let classes = object
|
||||
.get_type()
|
||||
.getattr("__mro__")?
|
||||
.cast_into::<PyTuple>()?
|
||||
.iter()
|
||||
.map(|class| {
|
||||
let class = class.cast_into::<PyType>()?;
|
||||
let attributes = class
|
||||
.getattr("__dict__")?
|
||||
.call_method0("items")?
|
||||
.try_iter()?
|
||||
.map(|item| item?.extract::<(String, Py<PyAny>)>())
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
Ok(ClassGuard {
|
||||
class: class.unbind(),
|
||||
attributes,
|
||||
})
|
||||
})
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
let guard = Self {
|
||||
reference: py
|
||||
.import("weakref")?
|
||||
.getattr("ref")?
|
||||
.call1((object,))?
|
||||
.unbind(),
|
||||
classes,
|
||||
config_names,
|
||||
config: Self::config(object, config_names)?,
|
||||
};
|
||||
if !guard.matches(py, object)? {
|
||||
return Err(PyTypeError::new_err(
|
||||
"native facade registration requires unmodified built-in methods",
|
||||
));
|
||||
}
|
||||
Ok(guard)
|
||||
}
|
||||
|
||||
fn config(object: &Bound<'_, PyAny>, names: &[&str]) -> PyResult<Vec<Value>> {
|
||||
names
|
||||
.iter()
|
||||
.map(|name| match object.getattr(*name) {
|
||||
Ok(value) => from_py(&value),
|
||||
Err(error)
|
||||
if error.is_instance_of::<pyo3::exceptions::PyAttributeError>(object.py()) =>
|
||||
{
|
||||
Ok(Value::Null)
|
||||
}
|
||||
Err(error) => Err(error),
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn matches(&self, py: Python<'_>, object: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
if !self.reference.bind(py).call0()?.is(object) {
|
||||
return Ok(false);
|
||||
}
|
||||
let mro = object
|
||||
.get_type()
|
||||
.getattr("__mro__")?
|
||||
.cast_into::<PyTuple>()?;
|
||||
if mro.len() != self.classes.len() {
|
||||
return Ok(false);
|
||||
}
|
||||
let instance = object.getattr("__dict__")?.cast_into::<PyDict>()?;
|
||||
for (class, expected) in mro.iter().zip(&self.classes) {
|
||||
if !class.is(expected.class.bind(py)) {
|
||||
return Ok(false);
|
||||
}
|
||||
let attributes = class.getattr("__dict__")?;
|
||||
if attributes.len()? != expected.attributes.len() {
|
||||
return Ok(false);
|
||||
}
|
||||
for (name, value) in &expected.attributes {
|
||||
if (instance.contains(name)?
|
||||
&& !self.config_names.contains(&name.as_str())
|
||||
&& !INSTANCE_STATE.contains(&name.as_str()))
|
||||
|| !attributes.get_item(name)?.is(value.bind(py))
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(Self::config(object, self.config_names)? == self.config)
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.reference)?;
|
||||
for class in &self.classes {
|
||||
visit.call(&class.class)?;
|
||||
for (_, value) in &class.attributes {
|
||||
visit.call(value)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl RedisPoolGuard {
|
||||
fn capture(
|
||||
backend: &Bound<'_, PyAny>,
|
||||
client_name: &'static str,
|
||||
attributes: RedisPoolAttributes,
|
||||
) -> PyResult<Self> {
|
||||
let pool = backend.getattr(client_name)?.getattr(attributes.pool)?;
|
||||
Ok(Self {
|
||||
reference: pool.clone().unbind(),
|
||||
connection_class: pool.getattr(attributes.connection_class)?.unbind(),
|
||||
connection_kwargs: pool
|
||||
.getattr("connection_kwargs")?
|
||||
.call_method0("copy")?
|
||||
.unbind(),
|
||||
max_connections: attributes
|
||||
.max_connections
|
||||
.map(|name| pool.getattr(name)?.extract::<usize>())
|
||||
.transpose()?,
|
||||
client_name,
|
||||
attributes,
|
||||
})
|
||||
}
|
||||
|
||||
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
let pool = backend
|
||||
.getattr(self.client_name)?
|
||||
.getattr(self.attributes.pool)?;
|
||||
Ok(self.reference.bind(py).is(&pool)
|
||||
&& self
|
||||
.connection_class
|
||||
.bind(py)
|
||||
.is(&pool.getattr(self.attributes.connection_class)?)
|
||||
&& self.max_connections
|
||||
== self
|
||||
.attributes
|
||||
.max_connections
|
||||
.map(|name| pool.getattr(name)?.extract::<usize>())
|
||||
.transpose()?
|
||||
&& self
|
||||
.connection_kwargs
|
||||
.bind(py)
|
||||
.eq(pool.getattr("connection_kwargs")?)?)
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.reference)?;
|
||||
visit.call(&self.connection_class)?;
|
||||
visit.call(&self.connection_kwargs)
|
||||
}
|
||||
}
|
||||
|
||||
impl DiskStoreGuard {
|
||||
fn capture(backend: &Bound<'_, PyAny>) -> PyResult<Self> {
|
||||
let store = backend.getattr("disk_cache")?;
|
||||
Ok(Self {
|
||||
reference: store.clone().unbind(),
|
||||
directory: store.getattr("directory")?.extract()?,
|
||||
})
|
||||
}
|
||||
|
||||
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
let store = backend.getattr("disk_cache")?;
|
||||
Ok(self.reference.bind(py).is(&store)
|
||||
&& self.directory == store.getattr("directory")?.extract::<String>()?)
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.reference)
|
||||
}
|
||||
}
|
||||
|
||||
impl AzureBlobClientGuard {
|
||||
fn capture(backend: &Bound<'_, PyAny>) -> PyResult<Self> {
|
||||
let sync_client = backend.getattr("container_client")?;
|
||||
Ok(Self {
|
||||
url: sync_client.getattr("url")?.extract::<String>()?,
|
||||
container_name: sync_client.getattr("container_name")?.extract::<String>()?,
|
||||
sync_client: sync_client.unbind(),
|
||||
async_client: backend.getattr("async_container_client")?.unbind(),
|
||||
})
|
||||
}
|
||||
|
||||
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
let sync_client = backend.getattr("container_client")?;
|
||||
Ok(self.sync_client.bind(py).is(&sync_client)
|
||||
&& self
|
||||
.async_client
|
||||
.bind(py)
|
||||
.is(&backend.getattr("async_container_client")?)
|
||||
&& self.url == sync_client.getattr("url")?.extract::<String>()?
|
||||
&& self.container_name == sync_client.getattr("container_name")?.extract::<String>()?)
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.sync_client)?;
|
||||
visit.call(&self.async_client)
|
||||
}
|
||||
}
|
||||
|
||||
impl S3ClientGuard {
|
||||
fn capture(backend: &Bound<'_, PyAny>) -> PyResult<Self> {
|
||||
Ok(Self {
|
||||
reference: backend.getattr("s3_client")?.unbind(),
|
||||
})
|
||||
}
|
||||
|
||||
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
Ok(self.reference.bind(py).is(&backend.getattr("s3_client")?))
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.reference)
|
||||
}
|
||||
}
|
||||
|
||||
impl ConnectionGuard {
|
||||
fn capture(kind: &str, cluster: bool, backend: &Bound<'_, PyAny>) -> PyResult<Self> {
|
||||
Ok(match (kind, cluster) {
|
||||
("redis", false) => Self::RedisPool(RedisPoolGuard::capture(
|
||||
backend,
|
||||
"redis_client",
|
||||
STANDALONE_POOL,
|
||||
)?),
|
||||
("redis", true) => Self::RedisPool(RedisPoolGuard::capture(
|
||||
backend,
|
||||
"redis_client",
|
||||
CLUSTER_POOL,
|
||||
)?),
|
||||
("valkey-semantic", _) => Self::RedisPool(RedisPoolGuard::capture(
|
||||
backend,
|
||||
"sync_client",
|
||||
VALKEY_POOL,
|
||||
)?),
|
||||
("disk", _) => Self::None,
|
||||
("azure-blob", _) => Self::AzureBlob(AzureBlobClientGuard::capture(backend)?),
|
||||
("s3", _) => Self::S3(S3ClientGuard::capture(backend)?),
|
||||
_ => Self::None,
|
||||
})
|
||||
}
|
||||
|
||||
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
match self {
|
||||
Self::None => Ok(true),
|
||||
Self::RedisPool(guard) => guard.matches(py, backend),
|
||||
Self::AzureBlob(guard) => guard.matches(py, backend),
|
||||
Self::S3(guard) => guard.matches(py, backend),
|
||||
}
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
match self {
|
||||
Self::None => Ok(()),
|
||||
Self::RedisPool(guard) => guard.traverse(visit),
|
||||
Self::AzureBlob(guard) => guard.traverse(visit),
|
||||
Self::S3(guard) => guard.traverse(visit),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl FacadeGuard {
|
||||
pub(in crate::cache) fn capture(
|
||||
py: Python<'_>,
|
||||
facade: &Bound<'_, PyAny>,
|
||||
service: &NativeResponseCache,
|
||||
) -> PyResult<Self> {
|
||||
let identity = service.identity();
|
||||
let kind = identity.kind();
|
||||
let cache_type = py.import("litellm.caching.caching")?.getattr("Cache")?;
|
||||
if !facade.get_type().is(&cache_type) {
|
||||
return Err(PyTypeError::new_err(
|
||||
"only exact built-in Cache facades can be registered",
|
||||
));
|
||||
}
|
||||
let cluster = matches!(
|
||||
identity,
|
||||
BackendIdentity::Redis {
|
||||
topology: RedisTopology::Cluster { .. },
|
||||
..
|
||||
}
|
||||
);
|
||||
let (module, name) = match (kind, cluster) {
|
||||
("memory", _) => ("litellm.caching.in_memory_cache", "InMemoryCache"),
|
||||
("redis", false) => ("litellm.caching.redis_cache", "RedisCache"),
|
||||
("redis", true) => ("litellm.caching.redis_cluster_cache", "RedisClusterCache"),
|
||||
("redis_semantic", _) => ("litellm.caching.redis_semantic_cache", "RedisSemanticCache"),
|
||||
("qdrant_semantic", _) => (
|
||||
"litellm.caching.qdrant_semantic_cache",
|
||||
"QdrantSemanticCache",
|
||||
),
|
||||
("gcs", _) => ("litellm.caching.gcs_cache", "GCSCache"),
|
||||
("valkey-semantic", _) => (
|
||||
"litellm.caching.valkey_semantic_cache",
|
||||
"ValkeySemanticCache",
|
||||
),
|
||||
("disk", _) => ("litellm.caching.disk_cache", "DiskCache"),
|
||||
("azure-blob", _) => ("litellm.caching.azure_blob_cache", "AzureBlobCache"),
|
||||
("s3", _) => ("litellm.caching.s3_cache", "S3Cache"),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
let cache_kind = identity.cache_type();
|
||||
let backend = facade.getattr("cache")?;
|
||||
if facade.getattr("type")?.extract::<String>()? != cache_kind
|
||||
|| !backend.get_type().is(&py.import(module)?.getattr(name)?)
|
||||
{
|
||||
return Err(PyTypeError::new_err(
|
||||
"facade and native backend types must match",
|
||||
));
|
||||
}
|
||||
let config = match NativeCacheConfig::project(facade)? {
|
||||
CacheConfigProjection::Native(config) => *config,
|
||||
CacheConfigProjection::Unsupported(reason) => {
|
||||
return Err(PyTypeError::new_err(reason.message()));
|
||||
}
|
||||
};
|
||||
if let Some(message) = config.service_mismatch(service) {
|
||||
return Err(PyTypeError::new_err(message));
|
||||
}
|
||||
if kind == "redis_semantic"
|
||||
&& service
|
||||
.embedder_object()
|
||||
.is_none_or(|embedder| !backend.is(embedder.bind(py)))
|
||||
{
|
||||
return Err(PyTypeError::new_err(
|
||||
"facade backend must be the native embedder",
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
outer: ObjectGuard::capture(
|
||||
py,
|
||||
facade,
|
||||
&[
|
||||
"type",
|
||||
"mode",
|
||||
"ttl",
|
||||
"namespace",
|
||||
"supported_call_types",
|
||||
"redis_flush_size",
|
||||
"semantic_cache_scope",
|
||||
],
|
||||
)?,
|
||||
backend: ObjectGuard::capture(
|
||||
py,
|
||||
&backend,
|
||||
&[
|
||||
"namespace",
|
||||
"default_ttl",
|
||||
"max_size_in_memory",
|
||||
"max_size_per_item",
|
||||
"redis_kwargs",
|
||||
"redis_flush_size",
|
||||
"similarity_threshold",
|
||||
"distance_threshold",
|
||||
"embedding_model",
|
||||
"embedding_max_input_tokens",
|
||||
"embedding_timeout",
|
||||
"qdrant_api_base",
|
||||
"qdrant_api_key",
|
||||
"collection_name",
|
||||
"vector_size",
|
||||
"_index_name",
|
||||
"_redis_url",
|
||||
"similarity_threshold",
|
||||
"embedding_model",
|
||||
"index_name",
|
||||
"embedding_max_input_tokens",
|
||||
"embedding_timeout",
|
||||
"bucket_name",
|
||||
"key_prefix",
|
||||
"path_service_account",
|
||||
],
|
||||
)?,
|
||||
disk_store: (kind == "disk")
|
||||
.then(|| DiskStoreGuard::capture(&backend))
|
||||
.transpose()?,
|
||||
connection: ConnectionGuard::capture(kind, cluster, &backend)?,
|
||||
})
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn matches(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
facade: &Bound<'_, PyAny>,
|
||||
) -> PyResult<bool> {
|
||||
if !self.outer.matches(py, facade)? {
|
||||
return Ok(false);
|
||||
}
|
||||
let backend = facade.getattr("cache")?;
|
||||
if !self.backend.matches(py, &backend)? {
|
||||
return Ok(false);
|
||||
}
|
||||
if let Some(guard) = &self.disk_store
|
||||
&& !guard.matches(py, &backend)?
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
self.connection.matches(py, &backend)
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn traverse(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
self.outer.traverse(&visit)?;
|
||||
self.backend.traverse(&visit)?;
|
||||
if let Some(guard) = &self.disk_store {
|
||||
guard.traverse(&visit)?;
|
||||
}
|
||||
self.connection.traverse(&visit)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,510 +0,0 @@
|
|||
use std::{path::PathBuf, time::Duration};
|
||||
|
||||
use litellm_cache_redis::RedisTopology;
|
||||
|
||||
/// What makes a native backend the one a Python facade describes: the configuration a user can
|
||||
/// observe on the Python object, captured once so facade projection and native construction
|
||||
/// compare plain data instead of reaching into each backend type.
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub(in crate::cache) enum BackendIdentity {
|
||||
Memory {
|
||||
capacity: usize,
|
||||
max_entry_bytes: Option<usize>,
|
||||
default_ttl: Option<Duration>,
|
||||
},
|
||||
Redis {
|
||||
topology: RedisTopology,
|
||||
namespace: Option<String>,
|
||||
default_ttl: Option<Duration>,
|
||||
},
|
||||
S3 {
|
||||
bucket: String,
|
||||
key_prefix: String,
|
||||
region: String,
|
||||
endpoint: Option<String>,
|
||||
},
|
||||
Gcs {
|
||||
bucket_name: String,
|
||||
key_prefix: String,
|
||||
path_service_account: Option<String>,
|
||||
},
|
||||
Disk {
|
||||
directory: PathBuf,
|
||||
},
|
||||
AzureBlob {
|
||||
account_url: String,
|
||||
container: String,
|
||||
},
|
||||
RedisSemantic {
|
||||
index_name: String,
|
||||
/// The backend stores the threshold as `f32`; a facade's `f64` is compared at that width.
|
||||
similarity_threshold: f32,
|
||||
},
|
||||
ValkeySemantic {
|
||||
index_name: String,
|
||||
similarity_threshold: f64,
|
||||
},
|
||||
QdrantSemantic {
|
||||
collection_name: String,
|
||||
similarity_threshold: f64,
|
||||
vector_size: u64,
|
||||
embedding_model: String,
|
||||
},
|
||||
}
|
||||
|
||||
const TYPES: &str = "facade and native backend types must match";
|
||||
|
||||
impl BackendIdentity {
|
||||
pub(in crate::cache) fn kind(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Memory { .. } => "memory",
|
||||
Self::Redis { .. } => "redis",
|
||||
Self::S3 { .. } => "s3",
|
||||
Self::Gcs { .. } => "gcs",
|
||||
Self::ValkeySemantic { .. } => "valkey-semantic",
|
||||
Self::RedisSemantic { .. } => "redis_semantic",
|
||||
Self::QdrantSemantic { .. } => "qdrant_semantic",
|
||||
Self::Disk { .. } => "disk",
|
||||
Self::AzureBlob { .. } => "azure-blob",
|
||||
}
|
||||
}
|
||||
|
||||
/// The `LiteLLMCacheType` value a facade of this backend carries in `Cache.type`.
|
||||
pub(in crate::cache) fn cache_type(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Memory { .. } => "local",
|
||||
Self::Redis { .. } => "redis",
|
||||
Self::S3 { .. } => "s3",
|
||||
Self::Gcs { .. } => "gcs",
|
||||
Self::ValkeySemantic { .. } => "valkey-semantic",
|
||||
Self::RedisSemantic { .. } => "redis-semantic",
|
||||
Self::QdrantSemantic { .. } => "qdrant-semantic",
|
||||
Self::Disk { .. } => "disk",
|
||||
Self::AzureBlob { .. } => "azure-blob",
|
||||
}
|
||||
}
|
||||
|
||||
/// The first difference between the facade's configuration (`self`) and the native
|
||||
/// backend (`native`), in the order Python users see the attributes.
|
||||
pub(in crate::cache) fn mismatch(&self, native: &Self) -> Option<&'static str> {
|
||||
let mut differences: Vec<(bool, &'static str)> = Vec::new();
|
||||
let mut differs = |condition: bool, message: &'static str| {
|
||||
differences.push((condition, message));
|
||||
};
|
||||
match (self, native) {
|
||||
(
|
||||
Self::Memory {
|
||||
capacity,
|
||||
max_entry_bytes,
|
||||
default_ttl,
|
||||
},
|
||||
Self::Memory {
|
||||
capacity: native_capacity,
|
||||
max_entry_bytes: native_max_entry_bytes,
|
||||
default_ttl: native_default_ttl,
|
||||
},
|
||||
) => {
|
||||
differs(
|
||||
default_ttl != native_default_ttl,
|
||||
"facade and native backend default TTLs must match",
|
||||
);
|
||||
differs(
|
||||
capacity != native_capacity,
|
||||
"facade and native backend capacities must match",
|
||||
);
|
||||
differs(
|
||||
max_entry_bytes != native_max_entry_bytes,
|
||||
"facade and native backend item limits must match",
|
||||
);
|
||||
}
|
||||
(
|
||||
Self::Redis {
|
||||
topology,
|
||||
namespace,
|
||||
default_ttl,
|
||||
},
|
||||
Self::Redis {
|
||||
topology: native_topology,
|
||||
namespace: native_namespace,
|
||||
default_ttl: native_default_ttl,
|
||||
},
|
||||
) => {
|
||||
differs(
|
||||
default_ttl != native_default_ttl,
|
||||
"facade and native backend default TTLs must match",
|
||||
);
|
||||
differs(
|
||||
topology != native_topology,
|
||||
"facade and native backend topologies must match",
|
||||
);
|
||||
differs(
|
||||
namespace != native_namespace,
|
||||
"facade and native backend namespaces must match",
|
||||
);
|
||||
}
|
||||
(
|
||||
Self::S3 {
|
||||
bucket,
|
||||
key_prefix,
|
||||
region,
|
||||
endpoint,
|
||||
},
|
||||
Self::S3 {
|
||||
bucket: native_bucket,
|
||||
key_prefix: native_key_prefix,
|
||||
region: native_region,
|
||||
endpoint: native_endpoint,
|
||||
},
|
||||
) => {
|
||||
differs(
|
||||
bucket != native_bucket,
|
||||
"facade and native backend buckets must match",
|
||||
);
|
||||
differs(
|
||||
key_prefix != native_key_prefix,
|
||||
"facade and native backend key prefixes must match",
|
||||
);
|
||||
differs(
|
||||
region != native_region,
|
||||
"facade and native backend regions must match",
|
||||
);
|
||||
differs(
|
||||
endpoint != native_endpoint,
|
||||
"facade and native backend endpoints must match",
|
||||
);
|
||||
}
|
||||
(
|
||||
Self::Gcs {
|
||||
bucket_name,
|
||||
key_prefix,
|
||||
path_service_account,
|
||||
},
|
||||
Self::Gcs {
|
||||
bucket_name: native_bucket_name,
|
||||
key_prefix: native_key_prefix,
|
||||
path_service_account: native_path_service_account,
|
||||
},
|
||||
) => {
|
||||
differs(
|
||||
bucket_name != native_bucket_name,
|
||||
"facade and native backend buckets must match",
|
||||
);
|
||||
differs(
|
||||
key_prefix != native_key_prefix,
|
||||
"facade and native backend key prefixes must match",
|
||||
);
|
||||
differs(
|
||||
path_service_account != native_path_service_account,
|
||||
"facade and native backend credentials must match",
|
||||
);
|
||||
}
|
||||
(
|
||||
Self::Disk { directory },
|
||||
Self::Disk {
|
||||
directory: native_directory,
|
||||
},
|
||||
) => {
|
||||
let canonical = |path: &PathBuf| std::fs::canonicalize(path).ok();
|
||||
differs(
|
||||
canonical(directory) != canonical(native_directory),
|
||||
"facade and native backend directories must match",
|
||||
);
|
||||
}
|
||||
(
|
||||
Self::AzureBlob {
|
||||
account_url,
|
||||
container,
|
||||
},
|
||||
Self::AzureBlob {
|
||||
account_url: native_account_url,
|
||||
container: native_container,
|
||||
},
|
||||
) => {
|
||||
differs(
|
||||
account_url != native_account_url || container != native_container,
|
||||
"facade and native backend containers must match",
|
||||
);
|
||||
}
|
||||
(
|
||||
Self::RedisSemantic {
|
||||
index_name,
|
||||
similarity_threshold,
|
||||
},
|
||||
Self::RedisSemantic {
|
||||
index_name: native_index_name,
|
||||
similarity_threshold: native_similarity_threshold,
|
||||
},
|
||||
) => {
|
||||
differs(
|
||||
index_name != native_index_name,
|
||||
"facade and native backend index names must match",
|
||||
);
|
||||
differs(
|
||||
similarity_threshold != native_similarity_threshold,
|
||||
"facade and native backend similarity thresholds must match",
|
||||
);
|
||||
}
|
||||
(
|
||||
Self::ValkeySemantic {
|
||||
index_name,
|
||||
similarity_threshold,
|
||||
},
|
||||
Self::ValkeySemantic {
|
||||
index_name: native_index_name,
|
||||
similarity_threshold: native_similarity_threshold,
|
||||
},
|
||||
) => {
|
||||
differs(
|
||||
index_name != native_index_name
|
||||
|| similarity_threshold != native_similarity_threshold,
|
||||
"facade and native semantic settings must match",
|
||||
);
|
||||
}
|
||||
(
|
||||
Self::QdrantSemantic {
|
||||
collection_name,
|
||||
similarity_threshold,
|
||||
vector_size,
|
||||
embedding_model,
|
||||
},
|
||||
Self::QdrantSemantic {
|
||||
collection_name: native_collection_name,
|
||||
similarity_threshold: native_similarity_threshold,
|
||||
vector_size: native_vector_size,
|
||||
embedding_model: native_embedding_model,
|
||||
},
|
||||
) => {
|
||||
differs(
|
||||
collection_name != native_collection_name,
|
||||
"facade and native backend collections must match",
|
||||
);
|
||||
differs(
|
||||
similarity_threshold != native_similarity_threshold,
|
||||
"facade and native backend similarity thresholds must match",
|
||||
);
|
||||
differs(
|
||||
vector_size != native_vector_size,
|
||||
"facade and native backend vector sizes must match",
|
||||
);
|
||||
differs(
|
||||
embedding_model != native_embedding_model,
|
||||
"facade and native backend embedding models must match",
|
||||
);
|
||||
}
|
||||
_ => return Some(TYPES),
|
||||
}
|
||||
differences
|
||||
.into_iter()
|
||||
.find_map(|(condition, message)| condition.then_some(message))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_cache_redis::{RedisNode, RedisTopology};
|
||||
|
||||
use super::BackendIdentity;
|
||||
|
||||
fn memory() -> BackendIdentity {
|
||||
BackendIdentity::Memory {
|
||||
capacity: 200,
|
||||
max_entry_bytes: Some(1024),
|
||||
default_ttl: Some(Duration::from_secs(60)),
|
||||
}
|
||||
}
|
||||
|
||||
fn redis() -> BackendIdentity {
|
||||
BackendIdentity::Redis {
|
||||
topology: RedisTopology::Standalone,
|
||||
namespace: Some("team".into()),
|
||||
default_ttl: Some(Duration::from_secs(60)),
|
||||
}
|
||||
}
|
||||
|
||||
fn s3() -> BackendIdentity {
|
||||
BackendIdentity::S3 {
|
||||
bucket: "bucket".into(),
|
||||
key_prefix: "cache/".into(),
|
||||
region: "us-east-1".into(),
|
||||
endpoint: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn gcs() -> BackendIdentity {
|
||||
BackendIdentity::Gcs {
|
||||
bucket_name: "bucket".into(),
|
||||
key_prefix: "cache/".into(),
|
||||
path_service_account: Some("credentials.json".into()),
|
||||
}
|
||||
}
|
||||
|
||||
fn azure() -> BackendIdentity {
|
||||
BackendIdentity::AzureBlob {
|
||||
account_url: "https://account.blob.core.windows.net".into(),
|
||||
container: "cache".into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn redis_semantic() -> BackendIdentity {
|
||||
BackendIdentity::RedisSemantic {
|
||||
index_name: "idx".into(),
|
||||
similarity_threshold: 0.8,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn redis_semantic_thresholds_compare_at_backend_precision() {
|
||||
let facade = BackendIdentity::RedisSemantic {
|
||||
index_name: "idx".into(),
|
||||
similarity_threshold: 0.8_f64 as f32,
|
||||
};
|
||||
assert_eq!(facade.mismatch(&redis_semantic()), None);
|
||||
}
|
||||
|
||||
fn valkey_semantic() -> BackendIdentity {
|
||||
BackendIdentity::ValkeySemantic {
|
||||
index_name: "idx".into(),
|
||||
similarity_threshold: 0.8,
|
||||
}
|
||||
}
|
||||
|
||||
fn qdrant() -> BackendIdentity {
|
||||
BackendIdentity::QdrantSemantic {
|
||||
collection_name: "collection".into(),
|
||||
similarity_threshold: 0.8,
|
||||
vector_size: 1536,
|
||||
embedding_model: "text-embedding-3-small".into(),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identical_identities_have_no_mismatch() {
|
||||
for identity in [
|
||||
memory(),
|
||||
redis(),
|
||||
s3(),
|
||||
gcs(),
|
||||
azure(),
|
||||
redis_semantic(),
|
||||
valkey_semantic(),
|
||||
qdrant(),
|
||||
BackendIdentity::Disk {
|
||||
directory: std::env::temp_dir(),
|
||||
},
|
||||
] {
|
||||
assert_eq!(identity.mismatch(&identity), None, "{identity:?}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn different_kinds_report_a_type_mismatch() {
|
||||
assert_eq!(
|
||||
memory().mismatch(&redis()),
|
||||
Some("facade and native backend types must match")
|
||||
);
|
||||
assert_eq!(
|
||||
redis_semantic().mismatch(&valkey_semantic()),
|
||||
Some("facade and native backend types must match")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_first_differing_field_names_the_mismatch() {
|
||||
let BackendIdentity::Memory { capacity, .. } = memory() else {
|
||||
unreachable!()
|
||||
};
|
||||
assert_eq!(
|
||||
memory().mismatch(&BackendIdentity::Memory {
|
||||
capacity: capacity + 1,
|
||||
max_entry_bytes: Some(1),
|
||||
default_ttl: Some(Duration::from_secs(60)),
|
||||
}),
|
||||
Some("facade and native backend capacities must match")
|
||||
);
|
||||
assert_eq!(
|
||||
memory().mismatch(&BackendIdentity::Memory {
|
||||
capacity,
|
||||
max_entry_bytes: Some(1),
|
||||
default_ttl: Some(Duration::from_secs(61)),
|
||||
}),
|
||||
Some("facade and native backend default TTLs must match")
|
||||
);
|
||||
assert_eq!(
|
||||
redis().mismatch(&BackendIdentity::Redis {
|
||||
topology: RedisTopology::Cluster {
|
||||
startup_nodes: vec![RedisNode {
|
||||
host: "node".into(),
|
||||
port: 7000,
|
||||
}],
|
||||
},
|
||||
namespace: None,
|
||||
default_ttl: Some(Duration::from_secs(60)),
|
||||
}),
|
||||
Some("facade and native backend topologies must match")
|
||||
);
|
||||
assert_eq!(
|
||||
s3().mismatch(&BackendIdentity::S3 {
|
||||
bucket: "bucket".into(),
|
||||
key_prefix: "cache/".into(),
|
||||
region: "us-east-1".into(),
|
||||
endpoint: Some("http://localhost:9000".into()),
|
||||
}),
|
||||
Some("facade and native backend endpoints must match")
|
||||
);
|
||||
assert_eq!(
|
||||
gcs().mismatch(&BackendIdentity::Gcs {
|
||||
bucket_name: "bucket".into(),
|
||||
key_prefix: "cache/".into(),
|
||||
path_service_account: None,
|
||||
}),
|
||||
Some("facade and native backend credentials must match")
|
||||
);
|
||||
assert_eq!(
|
||||
azure().mismatch(&BackendIdentity::AzureBlob {
|
||||
account_url: "https://account.blob.core.windows.net".into(),
|
||||
container: "other".into(),
|
||||
}),
|
||||
Some("facade and native backend containers must match")
|
||||
);
|
||||
assert_eq!(
|
||||
valkey_semantic().mismatch(&BackendIdentity::ValkeySemantic {
|
||||
index_name: "idx".into(),
|
||||
similarity_threshold: 0.9,
|
||||
}),
|
||||
Some("facade and native semantic settings must match")
|
||||
);
|
||||
assert_eq!(
|
||||
qdrant().mismatch(&BackendIdentity::QdrantSemantic {
|
||||
collection_name: "collection".into(),
|
||||
similarity_threshold: 0.8,
|
||||
vector_size: 1536,
|
||||
embedding_model: "text-embedding-3-large".into(),
|
||||
}),
|
||||
Some("facade and native backend embedding models must match")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disk_directories_compare_canonically() {
|
||||
let directory = std::env::temp_dir();
|
||||
let mut indirect = directory.clone();
|
||||
indirect.push(".");
|
||||
assert_eq!(
|
||||
BackendIdentity::Disk {
|
||||
directory: directory.clone()
|
||||
}
|
||||
.mismatch(&BackendIdentity::Disk {
|
||||
directory: indirect
|
||||
}),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
BackendIdentity::Disk { directory }.mismatch(&BackendIdentity::Disk {
|
||||
directory: "/definitely/missing".into()
|
||||
}),
|
||||
Some("facade and native backend directories must match")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -1,258 +0,0 @@
|
|||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use litellm_cache::{ExactCacheContext, SemanticCacheContext};
|
||||
use litellm_cache_response::{CacheControls, CacheKeyField, CacheKeyInput, ResponseCacheRequest};
|
||||
use litellm_host_python::from_py;
|
||||
use pyo3::{exceptions::PyValueError, prelude::*};
|
||||
use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct RequestInput {
|
||||
key: CacheKeyInput,
|
||||
controls: Option<CacheControls>,
|
||||
ttl_seconds: Option<f64>,
|
||||
max_age_seconds: Option<f64>,
|
||||
messages: Option<Value>,
|
||||
input: Option<Value>,
|
||||
metadata: Option<Value>,
|
||||
litellm_metadata: Option<Value>,
|
||||
litellm_params: Option<Value>,
|
||||
scope: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(in crate::cache) struct NativeRequest {
|
||||
pub(super) key: CacheKeyInput,
|
||||
pub(super) controls: CacheControls,
|
||||
pub(super) ttl: Option<Duration>,
|
||||
pub(super) max_age: Option<Duration>,
|
||||
pub(super) messages: Option<Value>,
|
||||
pub(super) input: Option<Value>,
|
||||
pub(super) metadata: Option<Value>,
|
||||
pub(super) litellm_metadata: Option<Value>,
|
||||
pub(super) litellm_params: Option<Value>,
|
||||
pub(super) scope: Option<String>,
|
||||
}
|
||||
|
||||
impl NativeRequest {
|
||||
pub(super) fn exact(&self) -> ResponseCacheRequest<ExactCacheContext> {
|
||||
ResponseCacheRequest {
|
||||
key: self.key.clone(),
|
||||
controls: self.controls,
|
||||
context: ExactCacheContext { ttl: self.ttl },
|
||||
max_age: self.max_age,
|
||||
}
|
||||
}
|
||||
|
||||
/// The request as a semantic backend that keys on the caller's scope sees it.
|
||||
pub(super) fn semantic(&self) -> ResponseCacheRequest<SemanticCacheContext> {
|
||||
self.semantic_with(self.key.clone(), self.scope.clone())
|
||||
}
|
||||
|
||||
/// The request keyed the way Python's Valkey semantic cache keys it: prompt fields drop out
|
||||
/// and the tenant identifiers for `scope` join the key.
|
||||
pub(super) fn scoped_semantic(
|
||||
&self,
|
||||
scope: &str,
|
||||
) -> ResponseCacheRequest<SemanticCacheContext> {
|
||||
self.semantic_with(semantic_key(self, scope), Some(scope.to_owned()))
|
||||
}
|
||||
|
||||
fn semantic_with(
|
||||
&self,
|
||||
key: CacheKeyInput,
|
||||
scope: Option<String>,
|
||||
) -> ResponseCacheRequest<SemanticCacheContext> {
|
||||
ResponseCacheRequest {
|
||||
key,
|
||||
controls: self.controls,
|
||||
context: SemanticCacheContext {
|
||||
input: self.input.clone(),
|
||||
messages: self.messages.clone(),
|
||||
metadata: self.metadata.clone(),
|
||||
scope,
|
||||
ttl: self.ttl,
|
||||
},
|
||||
max_age: self.max_age,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn semantic_key(request: &NativeRequest, scope: &str) -> CacheKeyInput {
|
||||
let mut key = request.key.clone();
|
||||
if key.preset.is_some() {
|
||||
return key;
|
||||
}
|
||||
key.fields
|
||||
.retain(|field| !matches!(field.name.as_str(), "messages" | "prompt" | "input"));
|
||||
const TENANT: [&str; 3] = [
|
||||
"user_api_key",
|
||||
"user_api_key_team_id",
|
||||
"user_api_key_org_id",
|
||||
];
|
||||
let end_user = (scope == "end_user").then_some("user_api_key_end_user_id");
|
||||
for name in TENANT.into_iter().chain(end_user) {
|
||||
let sources = [
|
||||
request.metadata.as_ref(),
|
||||
request.litellm_metadata.as_ref(),
|
||||
request
|
||||
.litellm_params
|
||||
.as_ref()
|
||||
.and_then(|params| params.get("metadata")),
|
||||
request
|
||||
.litellm_params
|
||||
.as_ref()
|
||||
.and_then(|params| params.get("litellm_metadata")),
|
||||
];
|
||||
let Some(value) = sources.into_iter().flatten().find_map(|source| {
|
||||
source
|
||||
.as_object()
|
||||
.and_then(|values| values.get(name))
|
||||
.filter(|value| !value.is_null())
|
||||
}) else {
|
||||
continue;
|
||||
};
|
||||
let value = match value {
|
||||
Value::Null => continue,
|
||||
Value::String(text) => text.clone(),
|
||||
other => other.to_string(),
|
||||
};
|
||||
key.fields.push(CacheKeyField {
|
||||
name: name.to_owned(),
|
||||
value: Some(value),
|
||||
api_parameter: true,
|
||||
internal_parameter: false,
|
||||
});
|
||||
}
|
||||
key
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn request(value: &Bound<'_, PyAny>) -> PyResult<NativeRequest> {
|
||||
let input: RequestInput = from_py(value)?;
|
||||
request_input(input)
|
||||
}
|
||||
|
||||
fn request_input(input: RequestInput) -> PyResult<NativeRequest> {
|
||||
let controls = input.controls.unwrap_or_else(|| {
|
||||
ResponseCacheRequest::<ExactCacheContext>::new(input.key.clone()).controls
|
||||
});
|
||||
Ok(NativeRequest {
|
||||
key: input.key,
|
||||
controls,
|
||||
ttl: input.ttl_seconds.map(duration).transpose()?,
|
||||
max_age: input.max_age_seconds.map(duration).transpose()?,
|
||||
messages: input.messages,
|
||||
input: input.input,
|
||||
metadata: input.metadata,
|
||||
litellm_metadata: input.litellm_metadata,
|
||||
litellm_params: input.litellm_params,
|
||||
scope: input.scope,
|
||||
})
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn requests(value: &Bound<'_, PyAny>) -> PyResult<Vec<NativeRequest>> {
|
||||
from_py::<Vec<RequestInput>>(value)?
|
||||
.into_iter()
|
||||
.map(request_input)
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn duration(seconds: f64) -> PyResult<Duration> {
|
||||
Duration::try_from_secs_f64(seconds)
|
||||
.map_err(|_| PyValueError::new_err("cache durations must be finite and nonnegative"))
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn now() -> Duration {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_cache_response::{CacheControls, CacheKeyInput, cache_key};
|
||||
use serde_json::json;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use super::*;
|
||||
|
||||
fn native_request(key: CacheKeyInput, metadata: Value) -> NativeRequest {
|
||||
NativeRequest {
|
||||
key,
|
||||
controls: CacheControls::default(),
|
||||
ttl: None,
|
||||
max_age: None,
|
||||
messages: Some(json!([{"role": "user", "content": "prompt"}])),
|
||||
input: None,
|
||||
metadata: Some(metadata),
|
||||
litellm_metadata: None,
|
||||
litellm_params: None,
|
||||
scope: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn semantic_key_matches_python_scope_material() {
|
||||
let key = CacheKeyInput {
|
||||
fields: vec![
|
||||
CacheKeyField {
|
||||
name: "model".to_owned(),
|
||||
value: Some("gpt-4.1".to_owned()),
|
||||
api_parameter: true,
|
||||
internal_parameter: false,
|
||||
},
|
||||
CacheKeyField {
|
||||
name: "messages".to_owned(),
|
||||
value: Some("prompt".to_owned()),
|
||||
api_parameter: true,
|
||||
internal_parameter: false,
|
||||
},
|
||||
],
|
||||
..Default::default()
|
||||
};
|
||||
let request = native_request(
|
||||
key,
|
||||
json!({"user_api_key": "k1", "user_api_key_team_id": null}),
|
||||
);
|
||||
let expected = format!("{:x}", Sha256::digest(b"model: gpt-4.1user_api_key: k1"));
|
||||
assert_eq!(cache_key(&semantic_key(&request, "key")), expected);
|
||||
assert_eq!(cache_key(&request.scoped_semantic("key").key), expected);
|
||||
|
||||
let end_user_request = native_request(
|
||||
request.key.clone(),
|
||||
json!({"user_api_key": "k1", "user_api_key_end_user_id": "u1"}),
|
||||
);
|
||||
let expected = format!(
|
||||
"{:x}",
|
||||
Sha256::digest(b"model: gpt-4.1user_api_key: k1user_api_key_end_user_id: u1")
|
||||
);
|
||||
assert_eq!(
|
||||
cache_key(&semantic_key(&end_user_request, "end_user")),
|
||||
expected
|
||||
);
|
||||
|
||||
let preset_request = native_request(
|
||||
CacheKeyInput {
|
||||
preset: Some("preset-key".to_owned()),
|
||||
..Default::default()
|
||||
},
|
||||
json!({"user_api_key": "k1"}),
|
||||
);
|
||||
assert_eq!(
|
||||
semantic_key(&preset_request, "end_user").preset.as_deref(),
|
||||
Some("preset-key")
|
||||
);
|
||||
assert!(semantic_key(&preset_request, "end_user").fields.is_empty());
|
||||
assert_eq!(preset_request.semantic().context.scope, None);
|
||||
assert_eq!(
|
||||
preset_request
|
||||
.scoped_semantic("end_user")
|
||||
.context
|
||||
.scope
|
||||
.as_deref(),
|
||||
Some("end_user")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,211 +0,0 @@
|
|||
use crate::cache::cache_error;
|
||||
use crate::execution::run_async;
|
||||
use std::{collections::VecDeque, time::Duration};
|
||||
|
||||
use litellm_cache::Error;
|
||||
use litellm_host_python::{Execution, ExecutionBody, ExecutionStep};
|
||||
use pyo3::{
|
||||
PyTraverseError, PyVisit,
|
||||
exceptions::{PyException, PyRuntimeError},
|
||||
prelude::*,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
backend::{NativeResponseCache, SemanticReply},
|
||||
embedder::{PythonEmbedder, with_prepared_embedding},
|
||||
request::{NativeRequest, now},
|
||||
};
|
||||
|
||||
pub(super) enum SemanticOperation {
|
||||
Lookup(NativeRequest),
|
||||
/// A lookup that also reports the similarity, as `SemanticReply`.
|
||||
LookupSemantic(NativeRequest),
|
||||
Store(NativeRequest, Value),
|
||||
StoreBatch(VecDeque<(NativeRequest, Value)>),
|
||||
}
|
||||
|
||||
/// What an exception from the Python embedder means for the operation.
|
||||
#[derive(Clone, Copy)]
|
||||
pub(super) enum EmbeddingFailure {
|
||||
/// Raise the Python exception unchanged.
|
||||
Propagate,
|
||||
/// Treat the embedding as unavailable and let the backend report that.
|
||||
Unavailable,
|
||||
}
|
||||
|
||||
enum Phase {
|
||||
Start,
|
||||
AwaitingEmbedding,
|
||||
AwaitingBackend,
|
||||
}
|
||||
|
||||
/// Runs a semantic cache operation whose embedding comes from Python: await the Python
|
||||
/// embedder in the caller's event loop, seed the native backend with the vector, await the
|
||||
/// backend, and repeat for each entry of a batch.
|
||||
pub(super) struct SemanticExecution {
|
||||
service: NativeResponseCache,
|
||||
embedder: PythonEmbedder,
|
||||
failure: EmbeddingFailure,
|
||||
operation: SemanticOperation,
|
||||
pending: Option<(NativeRequest, Option<Value>)>,
|
||||
phase: Phase,
|
||||
now: Duration,
|
||||
}
|
||||
|
||||
impl SemanticExecution {
|
||||
pub(super) fn new(
|
||||
service: NativeResponseCache,
|
||||
embedder: PythonEmbedder,
|
||||
failure: EmbeddingFailure,
|
||||
operation: SemanticOperation,
|
||||
) -> Self {
|
||||
Self {
|
||||
service,
|
||||
embedder,
|
||||
failure,
|
||||
operation,
|
||||
pending: None,
|
||||
phase: Phase::Start,
|
||||
now: now(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Takes the next entry of the operation; `None` once a batch is exhausted.
|
||||
fn next_pending(&mut self) -> Option<(NativeRequest, Option<Value>)> {
|
||||
match &mut self.operation {
|
||||
SemanticOperation::Lookup(request) | SemanticOperation::LookupSemantic(request) => {
|
||||
Some((request.clone(), None))
|
||||
}
|
||||
SemanticOperation::Store(request, response) => {
|
||||
Some((request.clone(), Some(std::mem::take(response))))
|
||||
}
|
||||
SemanticOperation::StoreBatch(queue) => queue
|
||||
.pop_front()
|
||||
.map(|(request, response)| (request, Some(response))),
|
||||
}
|
||||
}
|
||||
|
||||
fn start(&mut self, py: Python<'_>) -> PyResult<ExecutionStep> {
|
||||
let Some(pending) = self.next_pending() else {
|
||||
return Ok(ExecutionStep::Return(py.None()));
|
||||
};
|
||||
let (request, response) = &pending;
|
||||
let enabled = match response {
|
||||
None => request.controls.reads(),
|
||||
Some(_) => request.controls.writes(),
|
||||
};
|
||||
let input = enabled
|
||||
.then(|| self.service.embedding_input(request))
|
||||
.flatten();
|
||||
self.pending = Some(pending);
|
||||
let Some(input) = input else {
|
||||
return self.backend_step(py, Err(Error::Unavailable));
|
||||
};
|
||||
let awaitable =
|
||||
self.embedder
|
||||
.async_embedding(py, &input.prompt, input.metadata.as_ref())?;
|
||||
self.phase = Phase::AwaitingEmbedding;
|
||||
Ok(ExecutionStep::Await(awaitable))
|
||||
}
|
||||
|
||||
fn embedded(&mut self, py: Python<'_>, result: PyResult<Py<PyAny>>) -> PyResult<ExecutionStep> {
|
||||
let seed = match result {
|
||||
Ok(vector) => {
|
||||
PythonEmbedder::extract(vector.into_bound(py)).map_err(|_| Error::Unavailable)
|
||||
}
|
||||
Err(error) => match self.embedding_failure() {
|
||||
EmbeddingFailure::Propagate => return Err(error),
|
||||
EmbeddingFailure::Unavailable if error.is_instance_of::<PyException>(py) => {
|
||||
Err(Error::Unavailable)
|
||||
}
|
||||
EmbeddingFailure::Unavailable => return Err(error),
|
||||
},
|
||||
};
|
||||
self.backend_step(py, seed)
|
||||
}
|
||||
|
||||
/// Python's semantic lookups catch embedding errors and stamp a similarity of `0.0`.
|
||||
fn embedding_failure(&self) -> EmbeddingFailure {
|
||||
match self.operation {
|
||||
SemanticOperation::LookupSemantic(_) => EmbeddingFailure::Unavailable,
|
||||
_ => self.failure,
|
||||
}
|
||||
}
|
||||
|
||||
fn backend_step(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
seed: Result<Vec<f32>, Error>,
|
||||
) -> PyResult<ExecutionStep> {
|
||||
self.phase = Phase::AwaitingBackend;
|
||||
let (request, response) = self.pending.take().ok_or_else(|| {
|
||||
PyRuntimeError::new_err("semantic execution resumed without a pending operation")
|
||||
})?;
|
||||
let service = self.service.clone();
|
||||
let now = self.now;
|
||||
let with_similarity = matches!(self.operation, SemanticOperation::LookupSemantic(_));
|
||||
let future = async move {
|
||||
match response {
|
||||
None if with_similarity => service
|
||||
.async_lookup_semantic(&request, now)
|
||||
.await
|
||||
.map(|lookup| Reply::Semantic(lookup.into())),
|
||||
None => service.async_lookup(&request, now).await.map(Reply::Plain),
|
||||
Some(response) => service
|
||||
.async_store(&request, response, now)
|
||||
.await
|
||||
.map(|_| Reply::Plain(None)),
|
||||
}
|
||||
};
|
||||
let awaitable = run_async(py, with_prepared_embedding(seed, future), cache_error)?;
|
||||
Ok(ExecutionStep::Await(awaitable.unbind()))
|
||||
}
|
||||
|
||||
fn resume_py(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
result: Option<PyResult<Py<PyAny>>>,
|
||||
) -> PyResult<ExecutionStep> {
|
||||
match (&self.phase, result) {
|
||||
(Phase::Start, None) => self.start(py),
|
||||
(Phase::AwaitingEmbedding, Some(result)) => self.embedded(py, result),
|
||||
(Phase::AwaitingBackend, Some(Err(error))) => Err(error),
|
||||
(Phase::AwaitingBackend, Some(Ok(value))) => {
|
||||
let more = matches!(
|
||||
&self.operation,
|
||||
SemanticOperation::StoreBatch(queue) if !queue.is_empty()
|
||||
);
|
||||
if more {
|
||||
self.phase = Phase::Start;
|
||||
return self.start(py);
|
||||
}
|
||||
Ok(ExecutionStep::Return(value))
|
||||
}
|
||||
_ => Err(PyRuntimeError::new_err(
|
||||
"invalid semantic cache execution state",
|
||||
)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
#[serde(untagged)]
|
||||
enum Reply {
|
||||
Plain(Option<Value>),
|
||||
Semantic(SemanticReply),
|
||||
}
|
||||
|
||||
impl ExecutionBody for SemanticExecution {
|
||||
fn resume(&mut self, result: Option<PyResult<Py<PyAny>>>) -> PyResult<ExecutionStep> {
|
||||
Python::attach(|py| self.resume_py(py, result))
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
self.embedder.traverse(visit)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn drive(py: Python<'_>, body: SemanticExecution) -> PyResult<Bound<'_, PyAny>> {
|
||||
Execution::new(body, crate::lifecycle::binding).into_coroutine(py)
|
||||
}
|
||||
|
|
@ -1,5 +1,8 @@
|
|||
use crate::cache::cache_error;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
use std::{
|
||||
sync::Arc,
|
||||
time::{Duration, SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use litellm_cache::{DeleteCache, DisconnectCache, PingCache};
|
||||
use litellm_host_python::{from_py, release_gil, to_py};
|
||||
|
|
@ -118,8 +121,8 @@ impl NativeCacheHandle {
|
|||
fn get(&self, py: Python<'_>, key: String) -> PyResult<Py<PyAny>> {
|
||||
self.check_process()?;
|
||||
let request = request(key, None)?;
|
||||
let value = release_gil(py, || self.backend.lookup(&request, super::request::now()))
|
||||
.map_err(cache_error)?;
|
||||
let value =
|
||||
release_gil(py, || self.backend.lookup(&request, now())).map_err(cache_error)?;
|
||||
to_py(py, &value)
|
||||
}
|
||||
|
||||
|
|
@ -134,10 +137,7 @@ impl NativeCacheHandle {
|
|||
self.check_process()?;
|
||||
let request = request(key, ttl)?;
|
||||
let value: Value = from_py(value)?;
|
||||
release_gil(py, || {
|
||||
self.backend.store(&request, value, super::request::now())
|
||||
})
|
||||
.map_err(cache_error)
|
||||
release_gil(py, || self.backend.store(&request, value, now())).map_err(cache_error)
|
||||
}
|
||||
|
||||
fn async_get<'py>(&self, py: Python<'py>, key: String) -> PyResult<Bound<'py, PyAny>> {
|
||||
|
|
@ -146,7 +146,7 @@ impl NativeCacheHandle {
|
|||
let backend = self.backend.clone();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { backend.async_lookup(&request, super::request::now()).await },
|
||||
async move { backend.async_lookup(&request, now()).await },
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
|
|
@ -165,11 +165,7 @@ impl NativeCacheHandle {
|
|||
let backend = self.backend.clone();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
backend
|
||||
.async_store(&request, value, super::request::now())
|
||||
.await
|
||||
},
|
||||
async move { backend.async_store(&request, value, now()).await },
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
|
|
@ -190,11 +186,7 @@ impl NativeCacheHandle {
|
|||
let backend = self.backend.clone();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
backend
|
||||
.async_store_batch(entries, super::request::now())
|
||||
.await
|
||||
},
|
||||
async move { backend.async_store_batch(entries, now()).await },
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
|
|
@ -345,3 +337,9 @@ pub(in crate::cache) fn configured(
|
|||
},
|
||||
))
|
||||
}
|
||||
|
||||
fn now() -> Duration {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -1,165 +0,0 @@
|
|||
use pyo3::{
|
||||
PyTraverseError, PyVisit,
|
||||
exceptions::{PyTypeError, PyValueError},
|
||||
prelude::*,
|
||||
types::{PyDict, PyList, PyTuple},
|
||||
};
|
||||
|
||||
use crate::cache::future::ready_none;
|
||||
|
||||
pub(in crate::cache) struct PythonCallback(Py<PyAny>);
|
||||
|
||||
impl PythonCallback {
|
||||
pub(in crate::cache) fn new(object: Py<PyAny>) -> Self {
|
||||
Self(object)
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn lookup<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
kwargs: Option<&Bound<'py, PyDict>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.0
|
||||
.bind(py)
|
||||
.call_method("get_cache", (), Some(callback_kwargs(kwargs)?))
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn async_lookup<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
kwargs: Option<&Bound<'py, PyDict>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.0
|
||||
.bind(py)
|
||||
.call_method("async_get_cache", (), Some(callback_kwargs(kwargs)?))
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn store(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
response: &Bound<'_, PyAny>,
|
||||
kwargs: Option<&Bound<'_, PyDict>>,
|
||||
) -> PyResult<()> {
|
||||
self.0
|
||||
.bind(py)
|
||||
.call_method("add_cache", (response,), Some(callback_kwargs(kwargs)?))
|
||||
.map(|_| ())
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn async_store<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
response: &Bound<'py, PyAny>,
|
||||
kwargs: Option<&Bound<'py, PyDict>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.0.bind(py).call_method(
|
||||
"async_add_cache",
|
||||
(response,),
|
||||
Some(callback_kwargs(kwargs)?),
|
||||
)
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn lookup_batch<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
requests: &Bound<'py, PyAny>,
|
||||
kwargs: Option<&Bound<'py, PyAny>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let results = PyList::empty(py);
|
||||
for kwargs in batch_callback_kwargs(requests, kwargs)? {
|
||||
results.append(
|
||||
self.0
|
||||
.bind(py)
|
||||
.call_method("get_cache", (), Some(&kwargs))?,
|
||||
)?;
|
||||
}
|
||||
Ok(results.into_any())
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn async_lookup_batch<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
requests: &Bound<'py, PyAny>,
|
||||
kwargs: Option<&Bound<'py, PyAny>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let awaitables = batch_callback_kwargs(requests, kwargs)?
|
||||
.iter()
|
||||
.map(|kwargs| {
|
||||
self.0
|
||||
.bind(py)
|
||||
.call_method("async_get_cache", (), Some(kwargs))
|
||||
})
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
py.import("asyncio")?
|
||||
.call_method1("gather", PyTuple::new(py, awaitables)?)
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn async_store_batch<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
result: Option<&Bound<'py, PyAny>>,
|
||||
kwargs: Option<&Bound<'py, PyDict>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let result = result.ok_or_else(|| {
|
||||
PyTypeError::new_err("Python cache callbacks require their original callback_result")
|
||||
})?;
|
||||
self.0.bind(py).call_method(
|
||||
"async_add_cache_pipeline",
|
||||
(result,),
|
||||
Some(callback_kwargs(kwargs)?),
|
||||
)
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn async_flush<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let object = self.0.bind(py);
|
||||
let backend = match object.getattr_opt("cache")? {
|
||||
Some(backend) if !backend.is_none() => backend,
|
||||
_ => object.clone(),
|
||||
};
|
||||
if backend.hasattr("async_flush_cache")? {
|
||||
return backend.call_method0("async_flush_cache");
|
||||
}
|
||||
backend.call_method0("flush_cache")?;
|
||||
ready_none(py)
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.0.bind(py).call_method0("ping")
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.0)
|
||||
}
|
||||
}
|
||||
|
||||
fn callback_kwargs<'a, 'py>(
|
||||
kwargs: Option<&'a Bound<'py, PyDict>>,
|
||||
) -> PyResult<&'a Bound<'py, PyDict>> {
|
||||
kwargs.ok_or_else(|| {
|
||||
PyTypeError::new_err("Python cache callbacks require their original callback_kwargs")
|
||||
})
|
||||
}
|
||||
|
||||
fn batch_callback_kwargs<'py>(
|
||||
requests: &Bound<'py, PyAny>,
|
||||
kwargs: Option<&Bound<'py, PyAny>>,
|
||||
) -> PyResult<Vec<Bound<'py, PyDict>>> {
|
||||
let kwargs = kwargs
|
||||
.ok_or_else(|| {
|
||||
PyTypeError::new_err(
|
||||
"Python cache callbacks require one original callback_kwargs mapping per request",
|
||||
)
|
||||
})?
|
||||
.try_iter()?
|
||||
.map(|item| Ok(item?.cast_into::<PyDict>()?))
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
if kwargs.len() != requests.len()? {
|
||||
return Err(PyValueError::new_err(
|
||||
"batch cache requests and callback_kwargs must have equal lengths",
|
||||
));
|
||||
}
|
||||
Ok(kwargs)
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -1,385 +0,0 @@
|
|||
use crate::execution::run_async;
|
||||
use litellm_cache_response::PartialHits;
|
||||
use litellm_host_python::{ExecutionStep, from_py, release_gil, to_py};
|
||||
use pyo3::{
|
||||
PyTraverseError, PyVisit,
|
||||
exceptions::{PyRuntimeError, PyValueError},
|
||||
prelude::*,
|
||||
types::PyDict,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
cache_error,
|
||||
future::{ready_none, ready_value},
|
||||
native::activation::activate,
|
||||
native::backend::{NativeResponseCache, SemanticReply},
|
||||
native::config::{CacheConfigProjection, NativeCacheConfig},
|
||||
native::request::{now, request, requests},
|
||||
python::PythonCallback,
|
||||
};
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
|
||||
pub(super) enum CacheBinding {
|
||||
Disabled,
|
||||
Native(NativeResponseCache),
|
||||
PythonCallback(PythonCallback),
|
||||
}
|
||||
|
||||
#[pyclass(frozen, name = "_ResponseCacheRuntime")]
|
||||
pub(crate) struct ResolvedCache {
|
||||
binding: CacheBinding,
|
||||
guard: Option<super::native::facade::FacadeGuard>,
|
||||
pid: u32,
|
||||
}
|
||||
|
||||
impl ResolvedCache {
|
||||
pub(super) fn new(binding: CacheBinding) -> Self {
|
||||
Self {
|
||||
binding,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn with_guard(mut self, guard: super::native::facade::FacadeGuard) -> Self {
|
||||
self.guard = Some(guard);
|
||||
self
|
||||
}
|
||||
|
||||
pub(super) fn native_service(&self) -> PyResult<Option<NativeResponseCache>> {
|
||||
self.check_process()?;
|
||||
Ok(match &self.binding {
|
||||
CacheBinding::Native(service) => Some(service.clone()),
|
||||
_ => None,
|
||||
})
|
||||
}
|
||||
|
||||
fn check_process(&self) -> PyResult<()> {
|
||||
if matches!(self.binding, CacheBinding::Native(_)) && self.pid != std::process::id() {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
"native cache bindings must be resolved again after fork",
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn lookup_step(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
input: &Bound<'_, PyAny>,
|
||||
kwargs: Option<&Bound<'_, PyDict>>,
|
||||
) -> PyResult<ExecutionStep> {
|
||||
self.check_process()?;
|
||||
let awaitable = match &self.binding {
|
||||
CacheBinding::Disabled => ready_none(py)?,
|
||||
CacheBinding::Native(service) => {
|
||||
let request = request(input)?;
|
||||
service.async_lookup_py(py, request)?
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => callback.async_lookup(py, kwargs)?,
|
||||
};
|
||||
Ok(ExecutionStep::Await(awaitable.unbind()))
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl ResolvedCache {
|
||||
#[staticmethod]
|
||||
pub(crate) fn from_selected(cache: &Bound<'_, PyAny>) -> PyResult<Self> {
|
||||
let py = cache.py();
|
||||
let binding = if cache.is_none() {
|
||||
CacheBinding::Disabled
|
||||
} else if let Some(runtime) = cache
|
||||
.getattr_opt("_native_cache")?
|
||||
.filter(|value| !value.is_none())
|
||||
{
|
||||
let resolved = runtime
|
||||
.getattr("native")?
|
||||
.extract::<PyRef<'_, ResolvedCache>>()?;
|
||||
match resolved.native_service()? {
|
||||
Some(service) => {
|
||||
if !resolved
|
||||
.guard
|
||||
.as_ref()
|
||||
.is_some_and(|guard| guard.matches(py, cache).unwrap_or(false))
|
||||
{
|
||||
return Err(RustBridgeDeclined::new_err(
|
||||
"native cache runtime no longer matches its facade",
|
||||
));
|
||||
}
|
||||
CacheBinding::Native(service)
|
||||
}
|
||||
None => CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind())),
|
||||
}
|
||||
} else {
|
||||
CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind()))
|
||||
};
|
||||
Ok(Self::new(binding))
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
fn from_cache(cache: &Bound<'_, PyAny>) -> PyResult<Self> {
|
||||
let config = match NativeCacheConfig::project(cache)? {
|
||||
CacheConfigProjection::Native(config) => *config,
|
||||
CacheConfigProjection::Unsupported(reason) => {
|
||||
return Err(RustBridgeDeclined::new_err(reason.message()));
|
||||
}
|
||||
};
|
||||
let backend = cache.getattr("cache")?;
|
||||
let service = activate(cache.py(), &backend, config)?;
|
||||
let resolved = Self::new(CacheBinding::Native(service.clone()));
|
||||
Ok(
|
||||
match super::native::facade::FacadeGuard::capture(cache.py(), cache, &service) {
|
||||
Ok(guard) => resolved.with_guard(guard),
|
||||
Err(_) => resolved,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
#[getter]
|
||||
fn kind(&self) -> &'static str {
|
||||
match self.binding {
|
||||
CacheBinding::Disabled => "disabled",
|
||||
CacheBinding::Native(_) => "native",
|
||||
CacheBinding::PythonCallback(_) => "python_callback",
|
||||
}
|
||||
}
|
||||
|
||||
#[pyo3(signature = (request, *, callback_kwargs=None))]
|
||||
fn lookup(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
request: &Bound<'_, PyAny>,
|
||||
callback_kwargs: Option<&Bound<'_, PyDict>>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Disabled => Ok(py.None()),
|
||||
CacheBinding::Native(service) => {
|
||||
let request = self::request(request)?;
|
||||
let service = service.clone();
|
||||
let response = release_gil(py, move || service.lookup(&request, now()))
|
||||
.map_err(cache_error)?;
|
||||
to_py(py, &response)
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => {
|
||||
callback.lookup(py, callback_kwargs).map(Bound::unbind)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// `(response, similarity)`: the similarity is `None` when the backend reports none.
|
||||
fn lookup_semantic(&self, py: Python<'_>, request: &Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Native(service) => {
|
||||
let request = self::request(request)?;
|
||||
let service = service.clone();
|
||||
let lookup = release_gil(py, move || service.lookup_semantic(&request, now()))
|
||||
.map_err(cache_error)?;
|
||||
to_py(py, &SemanticReply::from(lookup))
|
||||
}
|
||||
CacheBinding::Disabled => to_py(py, &SemanticReply(None, None)),
|
||||
CacheBinding::PythonCallback(_) => Err(PyRuntimeError::new_err(
|
||||
"semantic lookups require a native cache binding",
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn async_lookup_semantic<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
request: &Bound<'py, PyAny>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Native(service) => {
|
||||
service.async_lookup_semantic_py(py, self::request(request)?)
|
||||
}
|
||||
CacheBinding::Disabled => ready_value(py, &SemanticReply(None, None)),
|
||||
CacheBinding::PythonCallback(_) => Err(PyRuntimeError::new_err(
|
||||
"semantic lookups require a native cache binding",
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
#[pyo3(signature = (request, response, *, callback_kwargs=None))]
|
||||
fn store(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
request: &Bound<'_, PyAny>,
|
||||
response: &Bound<'_, PyAny>,
|
||||
callback_kwargs: Option<&Bound<'_, PyDict>>,
|
||||
) -> PyResult<()> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Disabled => Ok(()),
|
||||
CacheBinding::Native(service) => {
|
||||
let request = self::request(request)?;
|
||||
let response: Value = from_py(response)?;
|
||||
let service = service.clone();
|
||||
release_gil(py, move || service.store(&request, response, now()))
|
||||
.map_err(cache_error)
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => callback.store(py, response, callback_kwargs),
|
||||
}
|
||||
}
|
||||
|
||||
#[pyo3(signature = (requests, *, callback_kwargs=None))]
|
||||
fn lookup_batch(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
requests: &Bound<'_, PyAny>,
|
||||
callback_kwargs: Option<&Bound<'_, PyAny>>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Disabled => {
|
||||
let requests = self::requests(requests)?;
|
||||
to_py(py, &PartialHits::new(vec![None; requests.len()]))
|
||||
}
|
||||
CacheBinding::Native(service) => {
|
||||
let requests = self::requests(requests)?;
|
||||
let service = service.clone();
|
||||
let response = release_gil(py, move || service.lookup_batch(&requests, now()))
|
||||
.map_err(cache_error)?;
|
||||
to_py(py, &response)
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => callback
|
||||
.lookup_batch(py, requests, callback_kwargs)
|
||||
.map(Bound::unbind),
|
||||
}
|
||||
}
|
||||
|
||||
#[pyo3(signature = (request, *, callback_kwargs=None))]
|
||||
fn async_lookup<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
request: &Bound<'py, PyAny>,
|
||||
callback_kwargs: Option<&Bound<'py, PyDict>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let ExecutionStep::Await(awaitable) = self.lookup_step(py, request, callback_kwargs)?
|
||||
else {
|
||||
unreachable!()
|
||||
};
|
||||
Ok(awaitable.into_bound(py))
|
||||
}
|
||||
|
||||
#[pyo3(signature = (request, response, *, callback_kwargs=None))]
|
||||
fn async_store<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
request: &Bound<'py, PyAny>,
|
||||
response: &Bound<'py, PyAny>,
|
||||
callback_kwargs: Option<&Bound<'py, PyDict>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Disabled => ready_none(py),
|
||||
CacheBinding::Native(service) => {
|
||||
let request = self::request(request)?;
|
||||
let response: Value = from_py(response)?;
|
||||
service.async_store_py(py, request, response)
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => {
|
||||
callback.async_store(py, response, callback_kwargs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[pyo3(signature = (requests, *, callback_kwargs=None))]
|
||||
fn async_lookup_batch<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
requests: &Bound<'py, PyAny>,
|
||||
callback_kwargs: Option<&Bound<'py, PyAny>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Disabled => {
|
||||
let requests = self::requests(requests)?;
|
||||
ready_value(py, &PartialHits::new(vec![None; requests.len()]))
|
||||
}
|
||||
CacheBinding::Native(service) => {
|
||||
let requests = self::requests(requests)?;
|
||||
let service = service.clone();
|
||||
run_async(
|
||||
py,
|
||||
async move { service.async_lookup_batch(&requests, now()).await },
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => {
|
||||
callback.async_lookup_batch(py, requests, callback_kwargs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[pyo3(signature = (requests, responses, *, callback_result=None, callback_kwargs=None))]
|
||||
fn async_store_batch<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
requests: &Bound<'py, PyAny>,
|
||||
responses: &Bound<'py, PyAny>,
|
||||
callback_result: Option<&Bound<'py, PyAny>>,
|
||||
callback_kwargs: Option<&Bound<'py, PyDict>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Disabled => ready_none(py),
|
||||
CacheBinding::Native(service) => {
|
||||
let requests = self::requests(requests)?;
|
||||
let responses: Vec<Value> = from_py(responses)?;
|
||||
if requests.len() != responses.len() {
|
||||
return Err(PyValueError::new_err(
|
||||
"batch cache requests and responses must have equal lengths",
|
||||
));
|
||||
}
|
||||
let entries = requests.into_iter().zip(responses).collect();
|
||||
service.async_store_batch_py(py, entries)
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => {
|
||||
callback.async_store_batch(py, callback_result, callback_kwargs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn async_flush<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Disabled => ready_none(py),
|
||||
CacheBinding::Native(service) => {
|
||||
let service = service.clone();
|
||||
run_async(py, async move { service.async_flush().await }, cache_error)
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => callback.async_flush(py),
|
||||
}
|
||||
}
|
||||
|
||||
fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
match &self.binding {
|
||||
CacheBinding::Disabled => ready_none(py),
|
||||
CacheBinding::Native(service) => {
|
||||
let service = service.clone();
|
||||
run_async(
|
||||
py,
|
||||
async move { service.test_connection().await },
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
CacheBinding::PythonCallback(callback) => callback.ping(py),
|
||||
}
|
||||
}
|
||||
|
||||
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
if let CacheBinding::PythonCallback(callback) = &self.binding {
|
||||
callback.traverse(&visit)?;
|
||||
}
|
||||
if let Some(guard) = &self.guard {
|
||||
guard.traverse(visit)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
@ -17,7 +17,6 @@ mod tokenizer;
|
|||
|
||||
#[pymodule(gil_used = true)]
|
||||
mod _native {
|
||||
use crate::cache::ResolvedCache;
|
||||
#[cfg(feature = "panic-test")]
|
||||
#[pymodule_export]
|
||||
use crate::diagnostics::_panic_for_test;
|
||||
|
|
@ -64,7 +63,6 @@ mod _native {
|
|||
"NativeCacheHandle",
|
||||
py.get_type::<crate::cache::NativeCacheHandle>(),
|
||||
)?;
|
||||
dict.set_item("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())?;
|
||||
dict.set_item(
|
||||
"_SecretManagerRuntime",
|
||||
py.get_type::<crate::secrets::runtime::NativeSecretManager>(),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)}
|
||||
11
tests/test_litellm_rust/cache/conftest.py
vendored
11
tests/test_litellm_rust/cache/conftest.py
vendored
|
|
@ -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()
|
||||
|
|
|
|||
165
tests/test_litellm_rust/cache/test_azure_blob.py
vendored
165
tests/test_litellm_rust/cache/test_azure_blob.py
vendored
|
|
@ -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()
|
||||
115
tests/test_litellm_rust/cache/test_disk.py
vendored
115
tests/test_litellm_rust/cache/test_disk.py
vendored
|
|
@ -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],
|
||||
}
|
||||
356
tests/test_litellm_rust/cache/test_facade.py
vendored
356
tests/test_litellm_rust/cache/test_facade.py
vendored
|
|
@ -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
|
||||
46
tests/test_litellm_rust/cache/test_gcs.py
vendored
46
tests/test_litellm_rust/cache/test_gcs.py
vendored
|
|
@ -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))
|
||||
|
|
@ -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"}
|
||||
212
tests/test_litellm_rust/cache/test_redis.py
vendored
212
tests/test_litellm_rust/cache/test_redis.py
vendored
|
|
@ -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
|
||||
560
tests/test_litellm_rust/cache/test_redis_semantic.py
vendored
560
tests/test_litellm_rust/cache/test_redis_semantic.py
vendored
|
|
@ -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"}
|
||||
253
tests/test_litellm_rust/cache/test_rollout.py
vendored
253
tests/test_litellm_rust/cache/test_rollout.py
vendored
|
|
@ -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
|
||||
167
tests/test_litellm_rust/cache/test_s3.py
vendored
167
tests/test_litellm_rust/cache/test_s3.py
vendored
|
|
@ -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)
|
||||
|
|
@ -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"}
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -1,112 +0,0 @@
|
|||
"""In-process path-style S3 stub for native cache parity tests."""
|
||||
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from email.utils import parsedate_to_datetime
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from typing import Final
|
||||
from urllib.parse import unquote, urlsplit
|
||||
|
||||
_STORED_HEADERS: Final = (
|
||||
"cache-control",
|
||||
"content-type",
|
||||
"content-language",
|
||||
"content-disposition",
|
||||
"expires",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class S3Object:
|
||||
body: bytes
|
||||
headers: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
|
||||
class S3Stub:
|
||||
"""Minimal path-style S3 endpoint serving PUT and GET object operations."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._objects: dict[str, S3Object] = {}
|
||||
stub: Final = self
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
def _key(self) -> str:
|
||||
parts: Final = urlsplit(self.path).path.lstrip("/").split("/", 1)
|
||||
return unquote(parts[1]) if len(parts) == 2 else ""
|
||||
|
||||
def _read_body(self) -> bytes:
|
||||
transfer: Final = self.headers.get("transfer-encoding", "")
|
||||
if "chunked" not in transfer:
|
||||
return self.rfile.read(int(self.headers.get("content-length", 0)))
|
||||
chunks: Final = bytearray()
|
||||
while True:
|
||||
size = int(self.rfile.readline().split(b";")[0].strip(), 16)
|
||||
if size == 0:
|
||||
while self.rfile.readline().strip():
|
||||
pass
|
||||
return bytes(chunks)
|
||||
chunks.extend(self.rfile.read(size))
|
||||
self.rfile.readline()
|
||||
|
||||
def do_PUT(self) -> None:
|
||||
body: Final = self._read_body()
|
||||
headers: Final = {name: self.headers[name] for name in _STORED_HEADERS if name in self.headers}
|
||||
stub._objects = {**stub._objects, self._key(): S3Object(body=body, headers=headers)}
|
||||
self.send_response(200)
|
||||
self.send_header("ETag", '"stub"')
|
||||
self.send_header("Content-Length", "0")
|
||||
self.end_headers()
|
||||
|
||||
def do_HEAD(self) -> None:
|
||||
self._object(send_body=False)
|
||||
|
||||
def do_GET(self) -> None:
|
||||
self._object(send_body=True)
|
||||
|
||||
def _object(self, send_body: bool) -> None:
|
||||
entry: Final = stub._objects.get(self._key())
|
||||
if entry is None:
|
||||
self.send_response(404)
|
||||
self.send_header("Content-Type", "application/xml")
|
||||
body: Final = b'<?xml version="1.0" encoding="UTF-8"?><Error><Code>NoSuchKey</Code></Error>'
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
if send_body:
|
||||
self.wfile.write(body)
|
||||
return
|
||||
self.send_response(200)
|
||||
for name, value in entry.headers.items():
|
||||
self.send_header(name, value)
|
||||
self.send_header("ETag", '"stub"')
|
||||
self.send_header("Content-Length", str(len(entry.body)))
|
||||
self.end_headers()
|
||||
if send_body:
|
||||
self.wfile.write(entry.body)
|
||||
|
||||
def log_message(self, format: str, *args: object) -> None:
|
||||
pass
|
||||
|
||||
self._server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
|
||||
self._worker: Final = threading.Thread(target=self._server.serve_forever, daemon=True)
|
||||
self._worker.start()
|
||||
|
||||
@property
|
||||
def url(self) -> str:
|
||||
host, port = self._server.server_address[:2]
|
||||
return f"http://{host}:{port}"
|
||||
|
||||
@property
|
||||
def objects(self) -> dict[str, S3Object]:
|
||||
return self._objects
|
||||
|
||||
def put_object(self, key: str, body: bytes, headers: dict[str, str] | None = None) -> None:
|
||||
self._objects = {**self._objects, key: S3Object(body=body, headers=headers or {})}
|
||||
|
||||
def expires(self, key: str) -> object:
|
||||
header: Final = self._objects[key].headers.get("expires")
|
||||
return parsedate_to_datetime(header) if header else None
|
||||
|
||||
def close(self) -> None:
|
||||
self._server.shutdown()
|
||||
self._server.server_close()
|
||||
self._worker.join(timeout=5)
|
||||
Loading…
Add table
Reference in a new issue