chore: merge main into litellm_cherry_pick_password_breach_reset

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
ryan 2026-09-21 21:54:35 +00:00
commit 7c8aed072f
266 changed files with 28211 additions and 3118 deletions

View file

@ -121,6 +121,10 @@ start_proxy() {
"LITELLM_MODEL_COST_MAP_URL=$INTEGRATION_UPSTREAM_URL/_cost_map"
"MODEL_COST_MAP_MIN_MODEL_COUNT=1"
"MODEL_COST_MAP_MAX_SHRINK_RATIO=0"
"GEMINI_API_BASE=$INTEGRATION_UPSTREAM_URL"
"ANTHROPIC_API_BASE=$INTEGRATION_UPSTREAM_URL"
"GEMINI_API_KEY=sk-scripted-provider"
"ANTHROPIC_API_KEY=sk-scripted-provider"
)
else
cost_map_env=("LITELLM_LOCAL_MODEL_COST_MAP=True")

View file

@ -205,7 +205,7 @@ def main(
native_module: Final = load_native_module(native_path)
native_module_loads: Final = native_module is not None
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
native_size_limit: Final = 25_000_000
native_size_limit: Final = 40_000_000
native_size_within_limit: Final = native_member.file_size <= native_size_limit
validations: Final = (
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),
@ -222,7 +222,7 @@ def main(
("Python extension entry point is present", extension_entry_point_present),
("Native module loads", native_module_loads),
("Production module omits the panic test hook", panic_test_hook_absent),
("Native extension does not exceed 25 MB", native_size_within_limit),
(f"Native extension does not exceed {native_size_limit / 1_000_000:.0f} MB", native_size_within_limit),
("Wheel contents are valid", not unexpected_members),
)
@ -267,7 +267,8 @@ def main(
),
(
not native_size_within_limit,
f"native extension exceeds 20 MB: {native_member.file_size / 1_000_000:.2f} MB",
f"native extension exceeds {native_size_limit / 1_000_000:.0f} MB: "
f"{native_member.file_size / 1_000_000:.2f} MB",
),
(bool(unexpected_members), f"wheel contains unexpected build artifacts: {', '.join(unexpected_members)}"),
)

View file

@ -130,7 +130,7 @@ jobs:
- name: Test secret manager feature combinations
run: |
cargo test -p litellm-auth-gcp --locked --no-default-features
for features in '' aws google aws,google; do
for features in '' aws google cyberark aws,google aws,google,cyberark; do
cargo test -p litellm-secrets --locked --no-default-features --features "$features"
done

View file

@ -307,6 +307,7 @@ For MCP OAuth, an upstream may advertise dynamic client registration but refuse
| [Deepgram (`deepgram`)](https://docs.litellm.ai/docs/providers/deepgram) | ✅ | ✅ | ✅ | | | ✅ | | | | |
| [DeepInfra (`deepinfra`)](https://docs.litellm.ai/docs/providers/deepinfra) | ✅ | ✅ | ✅ | | | | | | | |
| [Deepseek (`deepseek`)](https://docs.litellm.ai/docs/providers/deepseek) | ✅ | ✅ | ✅ | | | | | | | |
| [Eden AI (`edenai`)](https://docs.litellm.ai/docs/providers/edenai) | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | | | |
| [ElevenLabs (`elevenlabs`)](https://docs.litellm.ai/docs/providers/elevenlabs) | ✅ | ✅ | ✅ | | | ✅ | ✅ | | | |
| [Empower (`empower`)](https://docs.litellm.ai/docs/providers/empower) | ✅ | ✅ | ✅ | | | | | | | |
| [Fal AI (`fal_ai`)](https://docs.litellm.ai/docs/providers/fal_ai) | ✅ | ✅ | ✅ | | ✅ | | | | | |
@ -356,7 +357,7 @@ For MCP OAuth, an upstream may advertise dynamic client registration but refuse
| [Petals (`petals`)](https://docs.litellm.ai/docs/providers/petals) | ✅ | ✅ | ✅ | | | | | | | |
| [Pinstripes (`pinstripes`)](https://docs.litellm.ai/docs/providers/pinstripes) | ✅ | ✅ | ✅ | | | | | | | |
| [Predibase (`predibase`)](https://docs.litellm.ai/docs/providers/predibase) | ✅ | ✅ | ✅ | | | | | | | |
| [Qwen AI Platform (`qwen_ai_platform`)](https://docs.litellm.ai/docs/providers/qwencloud) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | ✅ |
| [Qianwen AI Platform (`qwen_ai_platform`)](https://docs.litellm.ai/docs/providers/qwencloud) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | ✅ |
| [QwenCloud (`qwencloud`)](https://docs.litellm.ai/docs/providers/qwencloud) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | ✅ |
| [Recraft (`recraft`)](https://docs.litellm.ai/docs/providers/recraft) | | | | | ✅ | | | | | |
| [Replicate (`replicate`)](https://docs.litellm.ai/docs/providers/replicate) | ✅ | ✅ | ✅ | | | | | | | |

View file

@ -0,0 +1 @@
ALTER TABLE "LiteLLM_PolicyAttachmentTable" ADD COLUMN IF NOT EXISTS "is_default" BOOLEAN NOT NULL DEFAULT false;

View file

@ -1421,6 +1421,7 @@ model LiteLLM_PolicyAttachmentTable {
models String[] @default([]) // Model names or patterns
tags String[] @default([]) // Tag patterns (e.g., ["healthcare", "prod-*"])
priority Int? // Explicit execution order
is_default Boolean @default(false) // Applied only when no non-default attachment matches
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt

158
litellm-rust/Cargo.lock generated
View file

@ -927,6 +927,12 @@ dependencies = [
"libc",
]
[[package]]
name = "crc16"
version = "0.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "338089f42c427b86394a5ee60ff321da23a5c89c9d89514c829687b26359fcff"
[[package]]
name = "crc32fast"
version = "1.5.1"
@ -2460,8 +2466,8 @@ dependencies = [
"rstest",
"serde",
"serde_json",
"sha2 0.10.9",
"thiserror 2.0.19",
"tokio",
]
[[package]]
@ -2479,12 +2485,29 @@ name = "litellm-cache-redis"
version = "0.1.0"
dependencies = [
"litellm-cache",
"r2d2",
"redis",
"redis-test",
"serde_json",
"tokio",
]
[[package]]
name = "litellm-cache-response"
version = "0.1.0"
dependencies = [
"litellm-cache",
"litellm-cache-memory",
"litellm-cache-redis",
"py_literal",
"redis",
"redis-test",
"serde",
"serde_json",
"sha2 0.10.9",
"tokio",
]
[[package]]
name = "litellm-callbacks-legacy-python"
version = "0.1.0"
@ -2648,6 +2671,10 @@ dependencies = [
"futures-util",
"litellm-auth",
"litellm-auth-gcp",
"litellm-cache",
"litellm-cache-memory",
"litellm-cache-redis",
"litellm-cache-response",
"litellm-callbacks-legacy-python",
"litellm-core",
"litellm-core-utils",
@ -2659,7 +2686,9 @@ dependencies = [
"pyo3",
"pyo3-async-runtimes",
"rstest",
"serde",
"serde_json",
"serde_with",
"tokio",
"tokio-tungstenite",
]
@ -2675,6 +2704,7 @@ dependencies = [
"jsonwebtoken",
"litellm-core-utils",
"litellm-secrets-aws",
"litellm-secrets-cyberark",
"litellm-secrets-google",
"litellm-secrets-types",
"moka",
@ -2709,6 +2739,26 @@ dependencies = [
"wiremock",
]
[[package]]
name = "litellm-secrets-cyberark"
version = "0.1.0"
dependencies = [
"base64 0.22.1",
"litellm-core-utils",
"litellm-secrets-types",
"moka",
"percent-encoding",
"reqwest 0.12.28",
"rstest",
"serde",
"serde_json",
"thiserror 2.0.19",
"tokio",
"tracing",
"veil",
"wiremock",
]
[[package]]
name = "litellm-secrets-google"
version = "0.1.0"
@ -2947,6 +2997,16 @@ dependencies = [
"minimal-lexical",
]
[[package]]
name = "num-bigint"
version = "0.4.8"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c89e69e7e0f03bea5ef08013795c25018e101932225a656383bd384495ecc367"
dependencies = [
"num-integer",
"num-traits",
]
[[package]]
name = "num-bigint"
version = "0.5.1"
@ -2957,6 +3017,15 @@ dependencies = [
"num-traits",
]
[[package]]
name = "num-complex"
version = "0.4.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "73f88a1307638156682bada9d7604135552957b7818057dcef22705b4d509495"
dependencies = [
"num-traits",
]
[[package]]
name = "num-conv"
version = "0.2.2"
@ -3130,6 +3199,48 @@ version = "2.3.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220"
[[package]]
name = "pest"
version = "2.9.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6d45aeb61b4bf818e12d4205f2466f8c4748f85f4fce0146d1c03d69d753f0ad"
dependencies = [
"memchr",
"ucd-trie",
]
[[package]]
name = "pest_derive"
version = "2.9.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "89cc5a242e25ed4e7704d0be240f2cfbe20a8c27e7e252d94835be93d92dc39f"
dependencies = [
"pest",
"pest_generator",
]
[[package]]
name = "pest_generator"
version = "2.9.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7abf21475cc3820fe4b2ca2dc2142902f67a02189f3b5b3a229f4febc01a43e5"
dependencies = [
"pest",
"pest_meta",
"proc-macro2",
"quote",
"syn 2.0.119",
]
[[package]]
name = "pest_meta"
version = "2.9.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "adba4db388f687393c18c51348d44a41d870ca9df71a2c98172ea3035dc6936e"
dependencies = [
"pest",
]
[[package]]
name = "pin-project"
version = "1.1.13"
@ -3304,6 +3415,19 @@ dependencies = [
"prost",
]
[[package]]
name = "py_literal"
version = "0.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "102df7a3d46db9d3891f178dcc826dc270a6746277a9ae6436f8d29fd490a8e1"
dependencies = [
"num-bigint 0.4.8",
"num-complex",
"num-traits",
"pest",
"pest_derive",
]
[[package]]
name = "pyo3"
version = "0.29.2"
@ -3469,6 +3593,17 @@ version = "6.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf"
[[package]]
name = "r2d2"
version = "0.8.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "51de85fb3fb6524929c8a2eb85e6b6d363de4e8c48f9e2c2eac4944abc181c93"
dependencies = [
"log",
"parking_lot",
"scheduled-thread-pool",
]
[[package]]
name = "rand"
version = "0.8.7"
@ -3602,9 +3737,13 @@ checksum = "2acbc41a996f7652b2ddd9dfd98cc4ff602cfd742ae35382f07f608405ab50ed"
dependencies = [
"arcstr",
"combine",
"crc16",
"itoa",
"num-bigint",
"num-bigint 0.5.1",
"percent-encoding",
"rand 0.10.2",
"rustls 0.23.42",
"rustls-native-certs",
"ryu",
"sha1_smol",
"socket2 0.6.5",
@ -4001,6 +4140,15 @@ dependencies = [
"windows-sys 0.61.2",
]
[[package]]
name = "scheduled-thread-pool"
version = "0.2.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3cbc66816425a074528352f5789333ecff06ca41b36b0b0efdfbb29edc391a19"
dependencies = [
"parking_lot",
]
[[package]]
name = "schemars"
version = "0.9.0"
@ -4954,6 +5102,12 @@ dependencies = [
"syn 2.0.119",
]
[[package]]
name = "ucd-trie"
version = "0.1.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2896d95c02a80c6d6a5d6e953d479f5ddf2dfdb6a244441010e373ac0fb88971"
[[package]]
name = "unarray"
version = "0.1.4"

View file

@ -22,12 +22,15 @@ litellm-secrets = { path = "crates/secrets" }
litellm-secrets-types = { path = "crates/secrets-types" }
litellm-secrets-aws = { path = "crates/secrets-aws" }
litellm-secrets-google = { path = "crates/secrets-google" }
litellm-secrets-cyberark = { path = "crates/secrets-cyberark" }
litellm-http = { path = "crates/http" }
litellm-llms = { path = "crates/llms" }
litellm-types = { path = "crates/types" }
litellm-core-utils = { path = "crates/core-utils" }
litellm-cache = { path = "crates/cache" }
litellm-cache-memory = { path = "crates/cache-memory" }
litellm-cache-redis = { path = "crates/cache-redis" }
litellm-cache-response = { path = "crates/cache-response" }
litellm-token-counter = { path = "crates/token-counter" }
litellm-token-counter-fast = { path = "crates/token-counter-fast" }
litellm-token-counter-huggingface = { path = "crates/token-counter-huggingface" }

View file

@ -7,8 +7,8 @@ repository.workspace = true
[dependencies]
litellm-cache.workspace = true
serde_json.workspace = true
[dev-dependencies]
serde_json.workspace = true
rstest.workspace = true
tokio.workspace = true

View file

@ -1,18 +1,20 @@
use std::cmp::Reverse;
use std::collections::{BinaryHeap, HashMap};
use std::sync::{Arc, Mutex};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use std::{
cmp::Reverse,
collections::{BinaryHeap, HashMap, HashSet},
hash::Hash,
sync::{Arc, Mutex},
time::{Duration, SystemTime, UNIX_EPOCH},
};
use litellm_cache::{
BaseCache, CacheConnectionResult, CacheConnectionStatus, CacheEntry, CacheFuture, CacheKwargs,
Error,
BaseCache, BatchCache, CacheConnectionResult, CacheConnectionStatus, ClaimCache, CounterCache,
DeleteCache, Error, ExactCacheContext, FlushCache, IncrementOperation, SetCache, TtlCache,
};
const DEFAULT_MAX_SIZE_IN_MEMORY: usize = 200;
const DEFAULT_TTL: Duration = Duration::from_secs(600);
type ValueMeasure<V> = Arc<dyn Fn(&V) -> Result<usize, Error> + Send + Sync>;
type ValueValidator<V> = Arc<dyn Fn(&V) -> Result<(), Error> + Send + Sync>;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CacheWrite {
@ -33,7 +35,6 @@ pub struct InMemoryCache<V: Clone> {
default_ttl: Duration,
max_entry_bytes: Option<usize>,
measure_value: Option<ValueMeasure<V>>,
validate_value: Option<ValueValidator<V>>,
now: Arc<dyn Fn() -> Duration + Send + Sync>,
}
@ -77,7 +78,6 @@ impl<V: Clone> InMemoryCache<V> {
default_ttl: default_ttl.unwrap_or(DEFAULT_TTL),
max_entry_bytes,
measure_value,
validate_value: None,
now: Arc::new(now),
}
}
@ -91,9 +91,6 @@ impl<V: Clone> InMemoryCache<V> {
if self.max_size_in_memory == 0 {
return Ok(CacheWrite::Disabled);
}
if let Some(validate) = &self.validate_value {
validate(&value)?;
}
if let (Some(limit), Some(measure)) = (self.max_entry_bytes, &self.measure_value)
&& measure(&value)? > limit
{
@ -101,15 +98,13 @@ impl<V: Clone> InMemoryCache<V> {
}
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
Self::evict(&mut state, self.max_size_in_memory, now);
let key = key.into();
state.values.insert(key.clone(), value);
Self::evict(&mut state, self.max_size_in_memory, now, &key);
let expiration = state.expirations.get(&key).copied();
if expiration.is_none_or(|expiration| expiration < now) {
let expiration = now + ttl.unwrap_or(self.default_ttl);
state.expirations.insert(key.clone(), expiration);
state.expiration_heap.push(Reverse((expiration, key)));
Self::set_expiration(&mut state, &key, now + ttl.unwrap_or(self.default_ttl));
}
state.values.insert(key, value);
Ok(CacheWrite::Stored)
}
@ -126,6 +121,14 @@ impl<V: Clone> InMemoryCache<V> {
Ok(state.values.get(key).cloned())
}
pub fn max_size_in_memory(&self) -> usize {
self.max_size_in_memory
}
pub fn max_entry_bytes(&self) -> Option<usize> {
self.max_entry_bytes
}
pub fn expires_at(&self, key: &str) -> Result<Option<Duration>, Error> {
Ok(self
.state
@ -136,6 +139,25 @@ impl<V: Clone> InMemoryCache<V> {
.copied())
}
pub async fn async_get_ttl(&self, key: &str) -> Result<Option<Duration>, Error> {
self.expires_at(key)
}
pub async fn async_get_oldest_n_keys(&self, count: usize) -> Result<Vec<String>, Error> {
let state = self.state.lock().map_err(|_| Error::Unavailable)?;
let mut expirations = state
.expirations
.iter()
.map(|(key, expiration)| (key.clone(), *expiration))
.collect::<Vec<_>>();
expirations.sort_unstable_by_key(|(_, expiration)| *expiration);
Ok(expirations
.into_iter()
.take(count)
.map(|(key, _)| key)
.collect())
}
pub fn delete_cache(&self, key: &str) -> Result<(), Error> {
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
Self::remove(&mut state, key);
@ -150,7 +172,7 @@ impl<V: Clone> InMemoryCache<V> {
Ok(())
}
fn evict(state: &mut CacheState<V>, capacity: usize, now: Duration) {
fn evict(state: &mut CacheState<V>, capacity: usize, now: Duration, key: &str) {
while let Some(Reverse((expiration, key))) = state.expiration_heap.peek().cloned() {
if state.expirations.get(&key).copied() != Some(expiration) {
state.expiration_heap.pop();
@ -161,6 +183,9 @@ impl<V: Clone> InMemoryCache<V> {
break;
}
}
if state.values.contains_key(key) {
return;
}
while state.values.len() >= capacity {
let Some(Reverse((expiration, key))) = state.expiration_heap.pop() else {
break;
@ -171,84 +196,205 @@ impl<V: Clone> InMemoryCache<V> {
}
}
fn set_expiration(state: &mut CacheState<V>, key: &str, expiration: Duration) {
if state.expirations.get(key).copied() != Some(expiration) {
state.expirations.insert(key.into(), expiration);
state
.expiration_heap
.push(Reverse((expiration, key.into())));
}
}
fn remove(state: &mut CacheState<V>, key: &str) {
state.values.remove(key);
state.expirations.remove(key);
}
}
impl InMemoryCache<CacheEntry> {
pub fn response_cache(capacity: usize, ttl: Duration, max_entry_bytes: usize) -> Self {
Self::response_cache_with_clock(capacity, ttl, max_entry_bytes, || {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
})
}
pub fn response_cache_with_clock(
capacity: usize,
ttl: Duration,
max_entry_bytes: usize,
now: impl Fn() -> Duration + Send + Sync + 'static,
) -> Self {
let mut cache = Self::with_clock_and_size_measurement(
Some(capacity),
Some(ttl),
Some(max_entry_bytes),
Some(Arc::new(|entry: &CacheEntry| {
serde_json::to_vec(entry)
.map(|bytes| bytes.len())
.map_err(|_| Error::InvalidEntry)
})),
now,
impl<V> ClaimCache for InMemoryCache<V>
where
V: Clone + PartialEq + Send + Sync + 'static,
{
fn claim_cache(
&self,
key: &str,
candidate: V,
eligible: &[V],
context: ExactCacheContext,
) -> Result<V, Error> {
if self.max_size_in_memory == 0 {
return Ok(candidate);
}
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
Self::evict(&mut state, self.max_size_in_memory, now, key);
let existing = state
.values
.get(key)
.filter(|existing| eligible.is_empty() || eligible.contains(existing))
.cloned();
if let Some(existing) = &existing
&& eligible.is_empty()
&& *existing != candidate
{
return Ok(existing.clone());
}
let winner = existing.unwrap_or(candidate);
Self::set_expiration(
&mut state,
key,
now + self.get_ttl(&context).unwrap_or(self.default_ttl),
);
cache.validate_value = Some(Arc::new(|entry: &CacheEntry| {
entry
.timestamp
.is_finite()
.then_some(())
.ok_or(Error::InvalidEntry)
}));
cache
state.values.insert(key.into(), winner.clone());
Ok(winner)
}
}
impl BaseCache for InMemoryCache<CacheEntry> {
type Value = CacheEntry;
impl CounterCache for InMemoryCache<f64> {
fn increment_cache(
&self,
key: &str,
amount: f64,
context: ExactCacheContext,
) -> Result<f64, Error> {
if self.max_size_in_memory == 0 {
return Ok(amount);
}
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
Self::evict(&mut state, self.max_size_in_memory, now, key);
let value = state.values.get(key).copied().unwrap_or_default() + amount;
if !state.expirations.contains_key(key) {
Self::set_expiration(
&mut state,
key,
now + self.get_ttl(&context).unwrap_or(self.default_ttl),
);
}
state.values.insert(key.into(), value);
Ok(value)
}
}
fn default_ttl(&self) -> Duration {
self.default_ttl
impl InMemoryCache<f64> {
pub async fn async_increment_pipeline(
&self,
operations: Vec<IncrementOperation>,
) -> Result<Vec<f64>, Error> {
operations
.into_iter()
.map(|operation| {
self.increment_cache(
&operation.key,
operation.amount,
ExactCacheContext { ttl: operation.ttl },
)
})
.collect()
}
}
impl<V: Clone + Send + Sync + 'static> BaseCache for InMemoryCache<V> {
type Value = V;
type Context = ExactCacheContext;
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
context.ttl.or(Some(self.default_ttl))
}
fn set_cache(&self, key: &str, value: Self::Value, kwargs: CacheKwargs) -> Result<(), Error> {
let ttl = self.get_ttl(&kwargs);
fn set_cache(
&self,
key: &str,
value: Self::Value,
context: &ExactCacheContext,
) -> Result<(), Error> {
let ttl = self.get_ttl(context).unwrap_or(self.default_ttl);
self.set_cache(key, value, Some(ttl)).map(|_| ())
}
fn get_cache(&self, key: &str, _: &CacheKwargs) -> Result<Option<Self::Value>, Error> {
fn get_cache(&self, key: &str, _: &ExactCacheContext) -> Result<Option<Self::Value>, Error> {
self.get_cache(key)
}
fn delete_cache(&self, key: &str) -> Result<(), Error> {
self.delete_cache(key)
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
fn flush_cache(&self) -> Result<(), Error> {
self.flush_cache()
}
fn disconnect(&self) -> CacheFuture<'_, ()> {
Box::pin(async { Ok(()) })
}
fn test_connection(&self) -> CacheFuture<'_, CacheConnectionResult> {
Box::pin(async {
Ok(CacheConnectionResult {
status: CacheConnectionStatus::Success,
message: "In-memory cache connection test successful".into(),
error: None,
})
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
Ok(CacheConnectionResult {
status: CacheConnectionStatus::Success,
message: "In-memory cache connection test successful".into(),
error: None,
})
}
}
impl<V: Clone + Send + Sync + 'static> BatchCache for InMemoryCache<V> {}
impl<V: Clone + Send + Sync + 'static> DeleteCache for InMemoryCache<V> {
fn delete_cache(&self, key: &str) -> Result<(), Error> {
InMemoryCache::delete_cache(self, key)
}
}
impl<V: Clone + Send + Sync + 'static> FlushCache for InMemoryCache<V> {
fn flush_cache(&self) -> Result<(), Error> {
InMemoryCache::flush_cache(self)
}
}
impl<V: Clone + Send + Sync + 'static> TtlCache for InMemoryCache<V> {
async fn async_get_ttl(&self, key: &str) -> Result<Option<Duration>, Error> {
InMemoryCache::async_get_ttl(self, key).await
}
}
impl<T> SetCache for InMemoryCache<HashSet<T>>
where
T: Clone + Eq + Hash + Send + Sync + 'static,
{
type SetValue = T;
type SetResult = Vec<T>;
async fn async_set_cache_sadd(
&self,
key: &str,
values: Vec<Self::SetValue>,
ttl: Option<Duration>,
) -> Result<Self::SetResult, Error> {
if self.max_size_in_memory == 0 {
return Ok(values);
}
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
Self::evict(&mut state, self.max_size_in_memory, now, key);
let mut stored = state.values.get(key).cloned().unwrap_or_default();
stored.extend(values.iter().cloned());
if let (Some(limit), Some(measure)) = (self.max_entry_bytes, &self.measure_value)
&& measure(&stored)? > limit
{
return Ok(values);
}
if !state.expirations.contains_key(key) {
Self::set_expiration(&mut state, key, now + ttl.unwrap_or(self.default_ttl));
}
state.values.insert(key.into(), stored);
Ok(values)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn repeated_increments_keep_one_heap_entry_per_expiration() {
let cache = InMemoryCache::<f64>::new(Some(4), None);
for _ in 0..100 {
cache
.increment_cache("counter", 1.0, ExactCacheContext::default())
.unwrap();
}
assert_eq!(cache.state.lock().unwrap().expiration_heap.len(), 1);
}
}

View file

@ -1,8 +1,16 @@
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use std::{
collections::HashSet,
sync::{
Arc,
atomic::{AtomicU64, Ordering},
},
time::Duration,
};
use litellm_cache::{BaseCache, CacheConnectionStatus, CacheEntry, Error};
use litellm_cache::{
BaseCache, CacheBackend, CacheConnectionStatus, ClaimCache, CounterCache, DeleteCache, Error,
ExactCacheContext, IncrementOperation, SetCache, get_cache, set_cache,
};
use litellm_cache_memory::{CacheWrite, InMemoryCache};
use rstest::{fixture, rstest};
@ -84,66 +92,49 @@ fn capacity_evicts_earliest_and_ignores_stale_heap_entries(clock: Arc<AtomicU64>
}
#[test]
fn disabled_size_limited_and_synchronized_response_writes_are_observable() {
let disabled = InMemoryCache::<CacheEntry>::response_cache(0, Duration::from_secs(60), 80);
fn disabled_size_limited_and_validated_writes_are_observable() {
let cache = |capacity| {
InMemoryCache::with_clock_and_size_measurement(
Some(capacity),
Some(Duration::from_secs(60)),
Some(4),
Some(Arc::new(|value: &String| {
if value.is_empty() {
return Err(Error::InvalidEntry);
}
Ok(value.len())
})),
|| Duration::from_secs(100),
)
};
let disabled = cache(0);
assert_eq!(
disabled
.set_cache(
"a",
CacheEntry {
timestamp: 1.0,
response: serde_json::json!("x")
},
None
)
.unwrap(),
disabled.set_cache("a", "x".into(), None).unwrap(),
CacheWrite::Disabled
);
let cache = InMemoryCache::<CacheEntry>::response_cache(2, Duration::from_secs(60), 80);
let cache = cache(2);
assert_eq!(
cache
.set_cache(
"large",
CacheEntry {
timestamp: 1.0,
response: serde_json::json!("x".repeat(100))
},
None
)
.unwrap(),
cache.set_cache("large", "oversized".into(), None).unwrap(),
CacheWrite::TooLarge
);
cache
.set_cache(
"small",
CacheEntry {
timestamp: 1.0,
response: serde_json::json!("ok"),
},
None,
)
.unwrap();
assert!(cache.get_cache("small").unwrap().is_some());
assert_eq!(cache.get_cache("large").unwrap(), None);
assert_eq!(
cache
.set_cache(
"invalid",
CacheEntry {
timestamp: f64::NAN,
response: serde_json::json!("bad"),
},
None,
)
.unwrap_err(),
Error::InvalidEntry
cache.set_cache("small", "ok".into(), None).unwrap(),
CacheWrite::Stored
);
assert_eq!(cache.get_cache("small").unwrap(), Some("ok".into()));
assert_eq!(
cache.set_cache("invalid", String::new(), None),
Err(Error::InvalidEntry)
);
assert_eq!(cache.get_cache("invalid").unwrap(), None);
cache.delete_cache("small").unwrap();
cache.flush_cache().unwrap();
assert_eq!(cache.get_cache("small").unwrap(), None);
}
#[tokio::test]
async fn connection_test_matches_python_result_contract() {
let cache = InMemoryCache::<CacheEntry>::default();
let cache = InMemoryCache::<String>::default();
let result = BaseCache::test_connection(&cache).await.unwrap();
assert_eq!(result.status, CacheConnectionStatus::Success);
assert_eq!(result.message, "In-memory cache connection test successful");
@ -156,3 +147,222 @@ async fn connection_test_matches_python_result_contract() {
})
);
}
#[tokio::test]
async fn generic_consumers_share_typed_values_and_honor_expiration() {
let clock = clock();
let cache: CacheBackend<InMemoryCache<String>> = Arc::new(cache(clock.clone(), 4));
let reader = Arc::clone(&cache);
let context = ExactCacheContext {
ttl: Some(Duration::from_secs(5)),
};
set_cache(cache.as_ref(), "sync", "first".into(), &context).unwrap();
assert_eq!(
get_cache(reader.as_ref(), "sync", &context).unwrap(),
Some("first".into())
);
cache
.batch_cache_write("async", "second".into(), context.clone())
.await
.unwrap();
cache
.async_set_cache_pipeline(vec![("batch".into(), "third".into())], context.clone())
.await
.unwrap();
drop(cache);
for (key, value) in [("sync", "first"), ("async", "second"), ("batch", "third")] {
assert_eq!(
reader.async_get_cache(key, &context).await.unwrap(),
Some(value.into())
);
}
reader.async_delete_cache("async").await.unwrap();
assert_eq!(
reader.async_get_cache("async", &context).await.unwrap(),
None
);
clock.store(106, Ordering::SeqCst);
assert_eq!(get_cache(reader.as_ref(), "sync", &context).unwrap(), None);
assert_eq!(
reader.async_get_cache("batch", &context).await.unwrap(),
None
);
}
#[test]
fn claims_are_atomic_and_refresh_eligible_winners() {
let clock = clock();
let cache = InMemoryCache::with_clock(Some(4), Some(Duration::from_secs(60)), {
let clock = clock.clone();
move || Duration::from_secs(clock.load(Ordering::SeqCst))
});
let context = ExactCacheContext {
ttl: Some(Duration::from_secs(10)),
};
assert_eq!(
cache
.claim_cache("affinity", "first".to_string(), &[], context.clone())
.unwrap(),
"first"
);
clock.store(103, Ordering::SeqCst);
assert_eq!(
cache
.claim_cache("affinity", "second".to_string(), &[], context.clone())
.unwrap(),
"first"
);
assert_eq!(
cache.expires_at("affinity").unwrap(),
Some(Duration::from_secs(110))
);
clock.store(105, Ordering::SeqCst);
assert_eq!(
cache
.claim_cache(
"affinity",
"second".to_string(),
&["first".to_string(), "second".to_string()],
context,
)
.unwrap(),
"first"
);
assert_eq!(
cache.expires_at("affinity").unwrap(),
Some(Duration::from_secs(115))
);
}
#[test]
fn counters_increment_under_one_lock() {
let cache = InMemoryCache::<f64>::default();
assert_eq!(
CounterCache::increment_cache(&cache, "counter", 1.5, ExactCacheContext::default())
.unwrap(),
1.5
);
assert_eq!(
CounterCache::increment_cache(&cache, "counter", 2.0, ExactCacheContext::default())
.unwrap(),
3.5
);
}
#[rstest]
fn rewriting_an_existing_key_at_capacity_keeps_other_entries(clock: Arc<AtomicU64>) {
let cache = cache(clock, 2);
cache
.set_cache("hot", "1".into(), Some(Duration::from_secs(10)))
.unwrap();
cache
.set_cache("cold", "2".into(), Some(Duration::from_secs(20)))
.unwrap();
cache.set_cache("cold", "3".into(), None).unwrap();
assert_eq!(cache.get_cache("hot").unwrap(), Some("1".into()));
assert_eq!(cache.get_cache("cold").unwrap(), Some("3".into()));
cache
.claim_cache("cold", "4".into(), &[], ExactCacheContext::default())
.unwrap();
assert_eq!(cache.get_cache("hot").unwrap(), Some("1".into()));
cache.set_cache("new", "5".into(), None).unwrap();
assert_eq!(cache.get_cache("hot").unwrap(), None);
assert_eq!(cache.get_cache("cold").unwrap(), Some("3".into()));
assert_eq!(cache.get_cache("new").unwrap(), Some("5".into()));
}
#[test]
fn incrementing_an_existing_counter_at_capacity_keeps_every_counter() {
let cache = InMemoryCache::<f64>::new(Some(2), None);
for key in ["a", "b", "a", "b"] {
cache
.increment_cache(key, 1.0, ExactCacheContext::default())
.unwrap();
}
assert_eq!(cache.get_cache("a").unwrap(), Some(2.0));
assert_eq!(cache.get_cache("b").unwrap(), Some(2.0));
}
#[test]
fn disabled_cache_does_not_retain_claims_or_counters() {
let claims = InMemoryCache::<String>::new(Some(0), None);
assert_eq!(
claims
.claim_cache("key", "first".into(), &[], ExactCacheContext::default())
.unwrap(),
"first"
);
assert_eq!(claims.get_cache("key").unwrap(), None);
let counters = InMemoryCache::<f64>::new(Some(0), None);
assert_eq!(
counters
.increment_cache("key", 2.0, ExactCacheContext::default())
.unwrap(),
2.0
);
assert_eq!(counters.get_cache("key").unwrap(), None);
}
#[tokio::test]
async fn ttl_and_oldest_key_operations_use_the_stored_expirations() {
let clock = Arc::new(AtomicU64::new(100));
let cache = cache(clock, 3);
cache
.set_cache("later", "2".into(), Some(Duration::from_secs(20)))
.unwrap();
cache
.set_cache("first", "1".into(), Some(Duration::from_secs(10)))
.unwrap();
assert_eq!(
cache.async_get_ttl("first").await.unwrap(),
Some(Duration::from_secs(110))
);
assert_eq!(cache.async_get_oldest_n_keys(1).await.unwrap(), ["first"]);
assert_eq!(cache.async_get_ttl("missing").await.unwrap(), None);
}
#[tokio::test]
async fn increment_pipeline_preserves_operation_order() {
let cache = InMemoryCache::<f64>::new(Some(3), None);
assert_eq!(
cache
.async_increment_pipeline(vec![
IncrementOperation {
key: "a".into(),
amount: 1.0,
ttl: Some(Duration::from_secs(10)),
},
IncrementOperation {
key: "a".into(),
amount: 2.0,
ttl: Some(Duration::from_secs(20)),
},
])
.await
.unwrap(),
[1.0, 3.0]
);
assert_eq!(cache.get_cache("a").unwrap(), Some(3.0));
}
#[tokio::test]
async fn set_capability_preserves_python_result_and_deduplicates_storage() {
let cache = InMemoryCache::<HashSet<String>>::new(None, None);
let inserted = vec!["a".into(), "a".into(), "b".into()];
assert_eq!(
cache
.async_set_cache_sadd("members", inserted.clone(), None)
.await
.unwrap(),
inserted
);
assert_eq!(
cache.get_cache("members").unwrap(),
Some(HashSet::from(["a".into(), "b".into()]))
);
}

View file

@ -7,9 +7,10 @@ repository.workspace = true
[dependencies]
litellm-cache.workspace = true
redis = "1.7.0"
serde_json.workspace = true
redis = { version = "1.7.0", features = ["cluster", "tls-rustls"] }
r2d2 = "0.8.10"
tokio.workspace = true
[dev-dependencies]
redis-test = "1.0.4"
serde_json.workspace = true

View file

@ -1,58 +1,201 @@
use std::sync::{Arc, Mutex, MutexGuard};
use std::time::Duration;
use std::{
sync::{Arc, Mutex},
time::Duration,
};
use litellm_cache::{
BaseCache, CacheConnectionResult, CacheConnectionStatus, CacheEntry, CacheFuture, CacheKwargs,
Error,
BaseCache, BatchCache, BatchEntry, CacheCodec, CacheConnectionResult, CacheConnectionStatus,
ClaimCache, CounterCache, DeleteCache, Error, ExactCacheContext, FlushCache,
};
use redis::Commands;
use crate::topology::RedisTopology;
mod connection;
mod operations;
pub(crate) use connection::ConnectionRef;
use connection::{ClusterConnectionManager, ConnectionManager};
pub use operations::{
RedisArg, RedisLpopOperation, RedisLpopResult, RedisRpushOperation, RedisScript,
};
const DEFAULT_TTL: Duration = Duration::from_secs(600);
const KEY_PREFIX: &str = "litellm-cache:";
const REDIS_TIMEOUT: Duration = Duration::from_secs(5);
const REDIS_POOL_SIZE: u32 = 16;
pub struct RedisCache<C = redis::Connection> {
connection: Arc<Mutex<C>>,
default_ttl: Duration,
const INCREMENT_SCRIPT: &str = concat!(
"local value = redis.call('INCRBYFLOAT', KEYS[1], ARGV[1]); ",
"if redis.call('TTL', KEYS[1]) == -1 then ",
"redis.call('EXPIRE', KEYS[1], ARGV[2]); end; return value"
);
const CLAIM_SCRIPT: &str = concat!(
"local current = redis.call('GET', KEYS[1]); ",
"if ARGV[1] == '' then if current ~= false and current ~= '' then return 0; end; ",
"elseif current ~= ARGV[1] then return 0; end; ",
"if ARGV[3] ~= '' then redis.call('SET', KEYS[1], ARGV[3], 'EX', ARGV[2]); ",
"elseif ARGV[4] == '1' then redis.call('EXPIRE', KEYS[1], ARGV[2]); end; return 1"
);
const CLAIM_ATTEMPTS: usize = 8;
enum Connections<C> {
Pool(r2d2::Pool<ConnectionManager>),
Cluster(r2d2::Pool<ClusterConnectionManager>),
Fixed(Mutex<C>),
}
impl RedisCache<redis::Connection> {
pub fn new(url: &str, default_ttl: Option<Duration>) -> Result<Self, Error> {
let client = redis::Client::open(url).map_err(|_| Error::Unavailable)?;
let connection = client.get_connection().map_err(|_| Error::Unavailable)?;
Ok(Self::with_connection(connection, default_ttl))
}
}
impl<C> RedisCache<C>
impl<C> Connections<C>
where
C: redis::ConnectionLike + Send + 'static,
{
fn with_connection(connection: C, default_ttl: Option<Duration>) -> Self {
Self {
connection: Arc::new(Mutex::new(connection)),
fn execute<T>(
&self,
operation: impl FnOnce(&mut ConnectionRef<'_>) -> Result<T, Error>,
) -> Result<T, Error> {
match self {
Self::Pool(pool) => {
let mut pooled = pool.get().map_err(|_| Error::Unavailable)?;
let result = operation(&mut ConnectionRef::Node(&mut pooled.connection));
pooled.failed = matches!(result, Err(Error::Unavailable));
result
}
Self::Cluster(pool) => {
let mut pooled = pool.get().map_err(|_| Error::Unavailable)?;
let result = operation(&mut ConnectionRef::Cluster(&mut pooled.connection));
pooled.failed = matches!(result, Err(Error::Unavailable));
result
}
Self::Fixed(connection) => {
let mut connection = connection.lock().map_err(|_| Error::Unavailable)?;
operation(&mut ConnectionRef::Node(&mut *connection))
}
}
}
}
pub struct RedisCache<S, C = redis::Connection> {
connections: Arc<Connections<C>>,
default_ttl: Duration,
codec: S,
namespace: Option<String>,
topology: RedisTopology,
}
impl<S: CacheCodec> RedisCache<S> {
pub fn new(url: &str, default_ttl: Option<Duration>, codec: S) -> Result<Self, Error> {
Self::connect(url, &RedisTopology::Standalone, default_ttl, codec)
}
pub fn connect(
url: &str,
topology: &RedisTopology,
default_ttl: Option<Duration>,
codec: S,
) -> Result<Self, Error> {
let connections = match topology {
RedisTopology::Standalone => Connections::Pool(pool(ConnectionManager::open(url)?)?),
RedisTopology::Cluster { startup_nodes } => {
Connections::Cluster(pool(ClusterConnectionManager::open(url, startup_nodes)?)?)
}
};
Ok(Self {
connections: Arc::new(connections),
default_ttl: default_ttl.unwrap_or(DEFAULT_TTL),
codec,
namespace: None,
topology: topology.clone(),
})
}
}
fn pool<M: r2d2::ManageConnection>(manager: M) -> Result<r2d2::Pool<M>, Error> {
r2d2::Pool::builder()
.max_size(REDIS_POOL_SIZE)
.min_idle(Some(0))
.connection_timeout(REDIS_TIMEOUT)
.test_on_check_out(false)
.build(manager)
.map_err(|_| Error::Unavailable)
}
impl<S, C> RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
pub fn with_connection(connection: C, default_ttl: Option<Duration>, codec: S) -> Self {
Self {
connections: Arc::new(Connections::Fixed(Mutex::new(connection))),
default_ttl: default_ttl.unwrap_or(DEFAULT_TTL),
codec,
namespace: None,
topology: RedisTopology::Standalone,
}
}
fn connection(&self) -> Result<MutexGuard<'_, C>, Error> {
self.connection.lock().map_err(|_| Error::Unavailable)
pub fn with_namespace(self, namespace: Option<String>) -> Self {
Self {
namespace: namespace.filter(|value| !value.is_empty()),
..self
}
}
fn namespaced_key(key: &str) -> String {
format!("{KEY_PREFIX}{key}")
pub fn namespace(&self) -> Option<&str> {
self.namespace.as_deref()
}
fn namespaced_pattern() -> &'static str {
const PATTERN: &str = "litellm-cache:*";
PATTERN
pub fn topology(&self) -> &RedisTopology {
&self.topology
}
fn encode(value: &CacheEntry) -> Result<Vec<u8>, Error> {
serde_json::to_vec(value).map_err(|_| Error::InvalidEntry)
fn namespaced_key(&self, key: &str) -> String {
namespaced_key(self.namespace.as_deref(), key)
}
fn decode(value: Vec<u8>) -> Result<CacheEntry, Error> {
serde_json::from_slice(&value).map_err(|_| Error::InvalidEntry)
fn namespaced_pattern(&self) -> Result<String, Error> {
let namespace = self.namespace.as_ref().ok_or(Error::UnscopedFlush)?;
let escaped: String = namespace
.chars()
.flat_map(|ch| {
if matches!(ch, '*' | '?' | '[' | ']' | '\\') {
vec!['\\', ch]
} else {
vec![ch]
}
})
.collect();
Ok(format!("{escaped}:*"))
}
fn flush_matching(connection: &mut ConnectionRef<'_>, pattern: &str) -> Result<(), Error> {
connection.scan(pattern, 1000, |connection, keys| {
if !keys.is_empty() {
connection
.del::<_, usize>(keys)
.map_err(|_| Error::Unavailable)?;
}
Ok(true)
})
}
fn decode_response(&self, value: redis::Value) -> Result<Option<S::Value>, Error> {
match value {
redis::Value::Nil => Ok(None),
redis::Value::BulkString(bytes) => self.codec.decode(&bytes).map(Some),
redis::Value::SimpleString(text) => self.codec.decode(text.as_bytes()).map(Some),
_ => Err(Error::InvalidEntry),
}
}
fn decode_batch_response(&self, value: redis::Value) -> Result<BatchEntry<S::Value>, Error> {
match self.decode_response(value) {
Ok(Some(value)) => Ok(BatchEntry::Hit(value)),
Ok(None) => Ok(BatchEntry::Miss),
Err(Error::InvalidEntry) => Ok(BatchEntry::Invalid),
Err(error) => Err(error),
}
}
fn ttl_seconds(ttl: Duration) -> u64 {
@ -61,196 +204,418 @@ where
.max(1)
}
fn run_blocking<T, F>(connection: Arc<Mutex<C>>, operation: F) -> CacheFuture<'static, T>
async fn run_blocking<T, F>(connections: Arc<Connections<C>>, operation: F) -> Result<T, Error>
where
T: Send + 'static,
F: FnOnce(&mut C) -> Result<T, Error> + Send + 'static,
F: FnOnce(&mut ConnectionRef<'_>) -> Result<T, Error> + Send + 'static,
{
Box::pin(async move {
tokio::task::spawn_blocking(move || {
let mut connection = connection.lock().map_err(|_| Error::Unavailable)?;
operation(&mut connection)
})
tokio::task::spawn_blocking(move || connections.execute(operation))
.await
.map_err(|_| Error::Unavailable)?
})
}
}
impl<C> BaseCache for RedisCache<C>
fn namespaced_key(namespace: Option<&str>, key: &str) -> String {
match namespace {
Some(namespace) if !key.starts_with(&format!("{namespace}:")) => {
format!("{namespace}:{key}")
}
_ => key.into(),
}
}
impl<S, C> BaseCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
type Value = CacheEntry;
type Value = S::Value;
type Context = ExactCacheContext;
fn default_ttl(&self) -> Duration {
self.default_ttl
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
context.ttl.or(Some(self.default_ttl))
}
fn set_cache(&self, key: &str, value: Self::Value, kwargs: CacheKwargs) -> Result<(), Error> {
let payload = Self::encode(&value)?;
let ttl = Self::ttl_seconds(self.get_ttl(&kwargs));
self.connection()?
.set_ex::<_, _, ()>(Self::namespaced_key(key), payload, ttl)
.map_err(|_| Error::Unavailable)
}
fn get_cache(&self, key: &str, _: &CacheKwargs) -> Result<Option<Self::Value>, Error> {
self.connection()?
.get::<_, Option<Vec<u8>>>(Self::namespaced_key(key))
.map_err(|_| Error::Unavailable)?
.map(Self::decode)
.transpose()
}
fn delete_cache(&self, key: &str) -> Result<(), Error> {
self.connection()?
.del::<_, ()>(Self::namespaced_key(key))
.map_err(|_| Error::Unavailable)
}
fn flush_cache(&self) -> Result<(), Error> {
let mut connection = self.connection()?;
let keys = connection
.scan_match(Self::namespaced_pattern())
.map_err(|_| Error::Unavailable)?
.collect::<redis::RedisResult<Vec<String>>>()
.map_err(|_| Error::Unavailable)?;
if keys.is_empty() {
return Ok(());
}
connection
.del::<_, usize>(keys)
.map(|_| ())
.map_err(|_| Error::Unavailable)
}
fn async_set_cache<'a>(
&'a self,
key: &'a str,
fn set_cache(
&self,
key: &str,
value: Self::Value,
kwargs: CacheKwargs,
) -> CacheFuture<'a, ()> {
let payload = Self::encode(&value);
let key = Self::namespaced_key(key);
let ttl = Self::ttl_seconds(self.get_ttl(&kwargs));
Self::run_blocking(Arc::clone(&self.connection), move |connection| {
context: &ExactCacheContext,
) -> Result<(), Error> {
let payload = self.codec.encode(&value)?;
let ttl = Self::ttl_seconds(self.get_ttl(context).unwrap_or(self.default_ttl));
let key = self.namespaced_key(key);
self.connections.execute(|connection| {
connection
.set_ex::<_, _, ()>(key, payload?, ttl)
.set_ex::<_, _, ()>(key, payload, ttl)
.map_err(|_| Error::Unavailable)
})
}
fn async_get_cache<'a>(
&'a self,
key: &'a str,
_: &'a CacheKwargs,
) -> CacheFuture<'a, Option<Self::Value>> {
let key = Self::namespaced_key(key);
Box::pin(async move {
Self::run_blocking(Arc::clone(&self.connection), move |connection| {
connection
.get::<_, Option<Vec<u8>>>(key)
.map_err(|_| Error::Unavailable)
})
.await?
.map(Self::decode)
.transpose()
})
fn get_cache(&self, key: &str, _: &ExactCacheContext) -> Result<Option<Self::Value>, Error> {
let key = self.namespaced_key(key);
let value = self.connections.execute(|connection| {
connection
.get::<_, redis::Value>(key)
.map_err(|_| Error::Unavailable)
})?;
self.decode_response(value)
}
fn async_set_cache_pipeline<'a>(
&'a self,
async fn async_set_cache(
&self,
key: &str,
value: Self::Value,
context: ExactCacheContext,
) -> Result<(), Error> {
let payload = self.codec.encode(&value)?;
let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
connection
.set_ex::<_, _, ()>(key, payload, ttl)
.map_err(|_| Error::Unavailable)
})
.await
}
async fn async_get_cache(
&self,
key: &str,
_: &ExactCacheContext,
) -> Result<Option<Self::Value>, Error> {
let key = self.namespaced_key(key);
let value = Self::run_blocking(Arc::clone(&self.connections), move |connection| {
connection
.get::<_, redis::Value>(key)
.map_err(|_| Error::Unavailable)
})
.await?;
self.decode_response(value)
}
async fn async_set_cache_pipeline(
&self,
cache_list: Vec<(String, Self::Value)>,
kwargs: CacheKwargs,
) -> CacheFuture<'a, ()> {
context: ExactCacheContext,
) -> Result<(), Error> {
let entries = cache_list
.into_iter()
.map(|(key, value)| {
Self::encode(&value).map(|payload| (Self::namespaced_key(&key), payload))
self.codec
.encode(&value)
.map(|payload| (self.namespaced_key(&key), payload))
})
.collect::<Result<Vec<_>, _>>();
let ttl = Self::ttl_seconds(self.get_ttl(&kwargs));
Self::run_blocking(Arc::clone(&self.connection), move |connection| {
for (key, payload) in entries? {
connection
.set_ex::<_, _, ()>(key, payload, ttl)
.map_err(|_| Error::Unavailable)?;
}
Ok(())
.collect::<Result<Vec<_>, _>>()?;
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
if entries.is_empty() {
return Ok(());
}
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
let commands = entries
.into_iter()
.map(|(key, payload)| {
let mut command = redis::cmd("SETEX");
command.arg(key).arg(ttl).arg(payload);
command
})
.collect();
connection.pipeline(commands).map(drop)
})
.await
}
fn async_delete_cache<'a>(&'a self, key: &'a str) -> CacheFuture<'a, ()> {
let key = Self::namespaced_key(key);
Self::run_blocking(Arc::clone(&self.connection), move |connection| {
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
match Self::run_blocking(Arc::clone(&self.connections), |connection| {
Ok(match connection.ping() {
Ok(_) => CacheConnectionResult {
status: CacheConnectionStatus::Success,
message: "Redis cache connection test successful".into(),
error: None,
},
Err(error) => CacheConnectionResult {
status: CacheConnectionStatus::Failed,
message: format!("Redis connection failed: {error}"),
error: Some(error.to_string()),
},
})
})
.await
{
Ok(result) => Ok(result),
Err(error) => Ok(CacheConnectionResult {
status: CacheConnectionStatus::Failed,
message: format!("Redis connection failed: {error}"),
error: Some(error.to_string()),
}),
}
}
}
impl<S, C> BatchCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
fn batch_get_cache(
&self,
keys: &[String],
_: &ExactCacheContext,
) -> Result<Vec<BatchEntry<Self::Value>>, Error> {
let keys = keys
.iter()
.map(|key| self.namespaced_key(key))
.collect::<Vec<_>>();
let values = self.connections.execute(|connection| {
redis::cmd("MGET")
.arg(keys)
.query::<Vec<redis::Value>>(connection)
.map_err(|_| Error::Unavailable)
})?;
values
.into_iter()
.map(|value| self.decode_batch_response(value))
.collect()
}
async fn async_batch_get_cache(
&self,
keys: Vec<String>,
_: ExactCacheContext,
) -> Result<Vec<BatchEntry<Self::Value>>, Error> {
let keys = keys
.iter()
.map(|key| self.namespaced_key(key))
.collect::<Vec<_>>();
let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("MGET")
.arg(keys)
.query::<Vec<redis::Value>>(connection)
.map_err(|_| Error::Unavailable)
})
.await?;
values
.into_iter()
.map(|value| self.decode_batch_response(value))
.collect()
}
}
impl<S, C> DeleteCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
fn delete_cache(&self, key: &str) -> Result<(), Error> {
let key = self.namespaced_key(key);
self.connections
.execute(|connection| connection.del::<_, ()>(key).map_err(|_| Error::Unavailable))
}
async fn async_delete_cache(&self, key: &str) -> Result<(), Error> {
let key = self.namespaced_key(key);
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
connection.del::<_, ()>(key).map_err(|_| Error::Unavailable)
})
.await
}
}
impl<S, C> FlushCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
fn flush_cache(&self) -> Result<(), Error> {
let pattern = self.namespaced_pattern()?;
self.connections
.execute(|connection| Self::flush_matching(connection, &pattern))
}
fn disconnect(&self) -> CacheFuture<'_, ()> {
Box::pin(async { Ok(()) })
}
fn test_connection(&self) -> CacheFuture<'_, CacheConnectionResult> {
Box::pin(async move {
Self::run_blocking(Arc::clone(&self.connection), |connection| {
redis::cmd("PING")
.query::<String>(connection)
.map_err(|_| Error::Unavailable)
})
.await?;
Ok(CacheConnectionResult {
status: CacheConnectionStatus::Success,
message: "Redis cache connection test successful".into(),
error: None,
})
async fn async_flush_cache(&self) -> Result<(), Error> {
let pattern = self.namespaced_pattern()?;
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
Self::flush_matching(connection, &pattern)
})
.await
}
}
impl<S, C> CounterCache for RedisCache<S, C>
where
S: CacheCodec<Value = f64>,
C: redis::ConnectionLike + Send + 'static,
{
fn increment_cache(
&self,
key: &str,
amount: f64,
context: ExactCacheContext,
) -> Result<f64, Error> {
let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
self.connections
.execute(|connection| increment(connection, key, amount, ttl))
}
async fn async_increment(
&self,
key: &str,
amount: f64,
context: ExactCacheContext,
) -> Result<f64, Error> {
let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
increment(connection, key, amount, ttl)
})
.await
}
}
fn increment(
connection: &mut ConnectionRef<'_>,
key: String,
amount: f64,
ttl: u64,
) -> Result<f64, Error> {
redis::cmd("EVAL")
.arg(INCREMENT_SCRIPT)
.arg(1)
.arg(key)
.arg(amount)
.arg(ttl)
.query(connection)
.map_err(|_| Error::Unavailable)
}
fn stored_bytes(value: redis::Value) -> Result<Option<Vec<u8>>, Error> {
match value {
redis::Value::Nil => Ok(None),
redis::Value::BulkString(bytes) => Ok(Some(bytes)),
redis::Value::SimpleString(text) => Ok(Some(text.into_bytes())),
_ => Err(Error::InvalidEntry),
}
}
/// Eligibility is decided on decoded values, so a pin written by another encoder (Python's
/// `json.dumps` spacing or key order) still matches. The write is a compare-and-set on the
/// bytes that decision was made on, retried when another claimant wins the race.
fn claim<S: CacheCodec>(
connection: &mut ConnectionRef<'_>,
codec: &S,
key: &str,
candidate: S::Value,
eligible: &[S::Value],
ttl: u64,
) -> Result<S::Value, Error>
where
S::Value: PartialEq,
{
let payload = codec.encode(&candidate)?;
if payload.is_empty() {
return Err(Error::InvalidEntry);
}
for _ in 0..CLAIM_ATTEMPTS {
let current = stored_bytes(
connection
.get::<_, redis::Value>(key)
.map_err(|_| Error::Unavailable)?,
)?
.filter(|bytes| !bytes.is_empty());
let existing = current
.as_deref()
.and_then(|bytes| codec.decode(bytes).ok())
.filter(|existing| eligible.is_empty() || eligible.contains(existing));
let refresh = existing
.as_ref()
.is_some_and(|existing| !eligible.is_empty() || *existing == candidate);
let write: &[u8] = if existing.is_some() { b"" } else { &payload };
let applied = redis::cmd("EVAL")
.arg(CLAIM_SCRIPT)
.arg(1)
.arg(key)
.arg(current.as_deref().unwrap_or_default())
.arg(ttl)
.arg(write)
.arg(u8::from(refresh))
.query::<bool>(connection)
.map_err(|_| Error::Unavailable)?;
if applied {
return Ok(existing.unwrap_or(candidate));
}
}
Err(Error::Unavailable)
}
impl<S, C> ClaimCache for RedisCache<S, C>
where
S: CacheCodec + Clone + 'static,
S::Value: PartialEq,
C: redis::ConnectionLike + Send + 'static,
{
fn claim_cache(
&self,
key: &str,
candidate: S::Value,
eligible: &[S::Value],
context: ExactCacheContext,
) -> Result<S::Value, Error> {
let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
self.connections
.execute(|connection| claim(connection, &self.codec, &key, candidate, eligible, ttl))
}
async fn async_claim_cache(
&self,
key: &str,
candidate: S::Value,
eligible: Vec<S::Value>,
context: ExactCacheContext,
) -> Result<S::Value, Error> {
let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
let codec = self.codec.clone();
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
claim(connection, &codec, &key, candidate, &eligible, ttl)
})
.await
}
}
#[cfg(test)]
mod tests {
use super::RedisCache;
use litellm_cache::{BaseCache, CacheEntry, CacheKwargs};
use redis_test::{MockCmd, MockRedisConnection};
use serde_json::json;
use std::time::Duration;
fn entry() -> CacheEntry {
CacheEntry {
timestamp: 123.0,
response: json!({"choices": [{"text": "cached"}]}),
}
}
use litellm_cache::{
BaseCache, CacheCodec, DeleteCache, ExactCacheContext, FlushCache, JsonCodec,
};
use redis_test::{MockCmd, MockRedisConnection};
use serde_json::json;
#[test]
fn cache_entries_round_trip_through_json() {
let entry = entry();
let encoded = RedisCache::<redis::Connection>::encode(&entry).unwrap();
assert_eq!(
RedisCache::<redis::Connection>::decode(encoded).unwrap(),
entry
);
}
use super::RedisCache;
#[test]
fn invalid_json_is_rejected() {
assert!(RedisCache::<redis::Connection>::decode(b"not json".to_vec()).is_err());
fn entry() -> serde_json::Value {
json!({"deployment": "model-a", "cooldown_seconds": 30})
}
#[test]
fn ttl_seconds_rounds_up_and_keeps_expiration_positive() {
assert_eq!(
RedisCache::<redis::Connection>::ttl_seconds(Duration::ZERO),
RedisCache::<JsonCodec<serde_json::Value>>::ttl_seconds(Duration::ZERO),
1
);
assert_eq!(
RedisCache::<redis::Connection>::ttl_seconds(Duration::from_millis(1500)),
RedisCache::<JsonCodec<serde_json::Value>>::ttl_seconds(Duration::from_millis(1500)),
2
);
assert_eq!(
RedisCache::<redis::Connection>::ttl_seconds(Duration::from_secs(15)),
RedisCache::<JsonCodec<serde_json::Value>>::ttl_seconds(Duration::from_secs(15)),
15
);
}
@ -258,7 +623,9 @@ mod tests {
#[test]
fn redis_commands_round_trip_entries_and_delete_only_namespaced_keys() {
let value = entry();
let payload = RedisCache::<redis::Connection>::encode(&value).unwrap();
let payload = JsonCodec::<serde_json::Value>::new()
.encode(&value)
.unwrap();
let connection = MockRedisConnection::new([
MockCmd::new(
redis::cmd("SETEX")
@ -271,13 +638,17 @@ mod tests {
MockCmd::new(redis::cmd("DEL").arg("litellm-cache:key"), Ok(1u32)),
])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None);
let cache =
RedisCache::with_connection(connection, None, JsonCodec::<serde_json::Value>::new())
.with_namespace(Some("litellm-cache".into()));
cache
.set_cache("key", value.clone(), CacheKwargs::default())
.set_cache("key", value.clone(), &ExactCacheContext::default())
.unwrap();
assert_eq!(
cache.get_cache("key", &CacheKwargs::default()).unwrap(),
cache
.get_cache("key", &ExactCacheContext::default())
.unwrap(),
Some(value)
);
cache.delete_cache("key").unwrap();
@ -290,13 +661,17 @@ mod tests {
redis::cmd("SCAN")
.cursor_arg(0)
.arg("MATCH")
.arg("litellm-cache:*"),
.arg("litellm-cache:*")
.arg("COUNT")
.arg(1000),
Ok(redis_test::redis_value!(["0", ["litellm-cache:key"]])),
),
MockCmd::new(redis::cmd("DEL").arg("litellm-cache:key"), Ok(1u32)),
])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None);
let cache =
RedisCache::with_connection(connection, None, JsonCodec::<serde_json::Value>::new())
.with_namespace(Some("litellm-cache".into()));
cache.flush_cache().unwrap();
}
@ -305,7 +680,9 @@ mod tests {
async fn test_connection_runs_ping_off_executor() {
let connection = MockRedisConnection::new([MockCmd::new(redis::cmd("PING"), Ok("PONG"))])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None);
let cache =
RedisCache::with_connection(connection, None, JsonCodec::<serde_json::Value>::new())
.with_namespace(Some("litellm-cache".into()));
assert_eq!(
cache.test_connection().await.unwrap().status,

View file

@ -0,0 +1,392 @@
use std::collections::HashMap;
use litellm_cache::Error;
use redis::{
ConnectionAddr, ConnectionInfo, ConnectionLike, IntoConnectionInfo,
cluster::{ClusterClient, ClusterClientBuilder, ClusterConnection, NodeAddress},
cluster_routing::{
MultipleNodeRoutingInfo, ResponsePolicy, RoutingInfo, SingleNodeRoutingInfo, Slot,
},
};
use super::REDIS_TIMEOUT;
use crate::topology::RedisNode;
pub(super) struct PooledConnection<C> {
pub(super) connection: C,
pub(super) failed: bool,
}
/// Pools connections without a checkout PING, which would double every operation's round trips.
/// A timed-out command leaves its reply on the socket while redis still reports the connection
/// open, so any connection whose operation failed is discarded instead of being reused.
pub(super) struct ConnectionManager(redis::Client);
impl ConnectionManager {
pub(super) fn open(url: &str) -> Result<Self, Error> {
redis::Client::open(url)
.map(Self)
.map_err(|_| Error::Unavailable)
}
}
impl r2d2::ManageConnection for ConnectionManager {
type Connection = PooledConnection<redis::Connection>;
type Error = redis::RedisError;
fn connect(&self) -> Result<Self::Connection, redis::RedisError> {
let connection = self.0.get_connection()?;
connection.set_read_timeout(Some(REDIS_TIMEOUT))?;
connection.set_write_timeout(Some(REDIS_TIMEOUT))?;
Ok(PooledConnection {
connection,
failed: false,
})
}
fn is_valid(&self, connection: &mut Self::Connection) -> Result<(), redis::RedisError> {
redis::cmd("PING").query::<String>(&mut connection.connection)?;
Ok(())
}
fn has_broken(&self, connection: &mut Self::Connection) -> bool {
connection.failed || !redis::ConnectionLike::is_open(&connection.connection)
}
}
pub(super) struct ClusterConnectionManager(ClusterClient);
impl ClusterConnectionManager {
pub(super) fn open(url: &str, startup_nodes: &[RedisNode]) -> Result<Self, Error> {
if startup_nodes.is_empty() {
return Err(Error::Unavailable);
}
let info = url.into_connection_info().map_err(|_| Error::Unavailable)?;
let nodes = startup_nodes
.iter()
.map(|node| node_info(&info, node))
.collect::<Result<Vec<_>, _>>()?;
ClusterClientBuilder::new(nodes)
.connection_timeout(REDIS_TIMEOUT)
.response_timeout(REDIS_TIMEOUT)
.build()
.map(Self)
.map_err(|_| Error::Unavailable)
}
}
fn node_info(info: &ConnectionInfo, node: &RedisNode) -> Result<ConnectionInfo, Error> {
let addr = match info.addr() {
ConnectionAddr::Tcp(..) => ConnectionAddr::Tcp(node.host.clone(), node.port),
ConnectionAddr::TcpTls {
insecure,
tls_params,
..
} => ConnectionAddr::TcpTls {
host: node.host.clone(),
port: node.port,
insecure: *insecure,
tls_params: tls_params.clone(),
},
_ => return Err(Error::Unavailable),
};
Ok(info.clone().set_addr(addr))
}
impl r2d2::ManageConnection for ClusterConnectionManager {
type Connection = PooledConnection<ClusterConnection>;
type Error = redis::RedisError;
fn connect(&self) -> Result<Self::Connection, redis::RedisError> {
let connection = self.0.get_connection()?;
connection.set_read_timeout(Some(REDIS_TIMEOUT))?;
connection.set_write_timeout(Some(REDIS_TIMEOUT))?;
Ok(PooledConnection {
connection,
failed: false,
})
}
fn is_valid(&self, connection: &mut Self::Connection) -> Result<(), redis::RedisError> {
redis::cmd("PING").query::<String>(&mut connection.connection)?;
Ok(())
}
fn has_broken(&self, connection: &mut Self::Connection) -> bool {
connection.failed || !redis::ConnectionLike::is_open(&connection.connection)
}
}
pub(crate) enum ConnectionRef<'a> {
Node(&'a mut dyn redis::ConnectionLike),
Cluster(&'a mut ClusterConnection),
}
impl redis::ConnectionLike for ConnectionRef<'_> {
fn req_packed_command(&mut self, cmd: &[u8]) -> redis::RedisResult<redis::Value> {
match self {
Self::Node(connection) => connection.req_packed_command(cmd),
Self::Cluster(connection) => connection.req_packed_command(cmd),
}
}
fn req_packed_commands(
&mut self,
cmd: &[u8],
offset: usize,
count: usize,
) -> redis::RedisResult<Vec<redis::Value>> {
match self {
Self::Node(connection) => connection.req_packed_commands(cmd, offset, count),
Self::Cluster(connection) => connection.req_packed_commands(cmd, offset, count),
}
}
fn get_db(&self) -> i64 {
match self {
Self::Node(connection) => connection.get_db(),
Self::Cluster(connection) => redis::ConnectionLike::get_db(*connection),
}
}
fn supports_pipelining(&self) -> bool {
match self {
Self::Node(connection) => connection.supports_pipelining(),
Self::Cluster(connection) => redis::ConnectionLike::supports_pipelining(*connection),
}
}
fn check_connection(&mut self) -> bool {
match self {
Self::Node(connection) => connection.check_connection(),
Self::Cluster(connection) => connection.check_connection(),
}
}
fn is_open(&self) -> bool {
match self {
Self::Node(connection) => connection.is_open(),
Self::Cluster(connection) => redis::ConnectionLike::is_open(*connection),
}
}
}
impl ConnectionRef<'_> {
pub(crate) fn pipeline(
&mut self,
commands: Vec<redis::Cmd>,
) -> Result<Vec<redis::Value>, Error> {
match self {
Self::Node(connection) => {
let mut pipeline = redis::pipe();
for command in &commands {
pipeline.add_command(command.clone());
}
pipeline
.query::<Vec<redis::Value>>(*connection)
.map_err(|_| Error::Unavailable)
}
Self::Cluster(connection) => {
let mut replies: Vec<Option<redis::Value>> = vec![None; commands.len()];
for indices in slot_groups(&commands).into_values() {
let mut pipeline = redis::pipe();
for index in &indices {
pipeline.add_command(commands[*index].clone());
}
let values = connection
.req_packed_commands(&pipeline.get_packed_pipeline(), 0, indices.len())
.map_err(|_| Error::Unavailable)?;
if values.len() != indices.len() {
return Err(Error::Unavailable);
}
for (index, value) in indices.into_iter().zip(values) {
replies[index] = Some(value);
}
}
replies
.into_iter()
.collect::<Option<Vec<_>>>()
.ok_or(Error::Unavailable)
}
}
}
pub(crate) fn scan(
&mut self,
pattern: &str,
count: usize,
mut visit: impl FnMut(&mut Self, Vec<String>) -> Result<bool, Error>,
) -> Result<(), Error> {
let pages = match self {
Self::Node(connection) => {
let page = scan_command(0, pattern, count)
.query::<ScanPage>(*connection)
.map_err(|_| Error::Unavailable)?;
vec![(None, page)]
}
Self::Cluster(connection) => connection
.route_command(
&scan_command(0, pattern, count),
RoutingInfo::MultiNode((
MultipleNodeRoutingInfo::AllMasters,
Some(ResponsePolicy::Special),
)),
)
.map_err(|_| Error::Unavailable)
.and_then(primary_pages)?
.into_iter()
.map(|(node, page)| (Some(node), page))
.collect(),
};
for (node, (mut cursor, mut keys)) in pages {
loop {
if !visit(self, keys)? {
return Ok(());
}
if cursor == 0 {
break;
}
(cursor, keys) = self.scan_page(node.as_ref(), cursor, pattern, count)?;
}
}
Ok(())
}
pub(crate) fn ping(&mut self) -> Result<bool, redis::RedisError> {
let command = redis::cmd("PING");
match self {
Self::Node(connection) => command
.query::<String>(*connection)
.map(|response| response == "PONG"),
Self::Cluster(connection) => connection
.route_command(
&command,
RoutingInfo::MultiNode((
MultipleNodeRoutingInfo::AllNodes,
Some(ResponsePolicy::AllSucceeded),
)),
)
.map(|_| true),
}
}
pub(crate) fn node_text(&mut self, command: &redis::Cmd) -> Result<String, Error> {
match self {
Self::Node(connection) => command.query(*connection).map_err(|_| Error::Unavailable),
Self::Cluster(connection) => {
let value = connection
.route_command(
command,
RoutingInfo::MultiNode((
MultipleNodeRoutingInfo::AllNodes,
Some(ResponsePolicy::Special),
)),
)
.map_err(|_| Error::Unavailable)?;
let redis::Value::Map(entries) = value else {
return Err(Error::Unavailable);
};
let mut replies = entries
.into_iter()
.map(|(node, reply)| {
Ok((
redis::from_redis_value::<String>(node)
.map_err(|_| Error::Unavailable)?,
redis::from_redis_value::<String>(reply)
.map_err(|_| Error::Unavailable)?,
))
})
.collect::<Result<Vec<(String, String)>, Error>>()?;
replies.sort();
Ok(replies
.into_iter()
.map(|(_, reply)| reply)
.collect::<Vec<_>>()
.join("\n"))
}
}
}
pub(crate) fn flushall(&mut self) -> Result<(), Error> {
let command = redis::cmd("FLUSHALL");
match self {
Self::Node(connection) => command.query(*connection).map_err(|_| Error::Unavailable),
Self::Cluster(connection) => connection
.route_command(
&command,
RoutingInfo::MultiNode((
MultipleNodeRoutingInfo::AllMasters,
Some(ResponsePolicy::AllSucceeded),
)),
)
.map(|_| ())
.map_err(|_| Error::Unavailable),
}
}
fn scan_page(
&mut self,
node: Option<&NodeAddress>,
cursor: u64,
pattern: &str,
count: usize,
) -> Result<ScanPage, Error> {
let command = scan_command(cursor, pattern, count);
match (self, node) {
(Self::Node(connection), None) => {
command.query(*connection).map_err(|_| Error::Unavailable)
}
(Self::Cluster(connection), Some(node)) => connection
.route_command(
&command,
RoutingInfo::SingleNode(SingleNodeRoutingInfo::ByAddress {
host: node.host().to_string(),
port: node.port(),
}),
)
.map_err(|_| Error::Unavailable)
.and_then(|value| redis::from_redis_value(value).map_err(|_| Error::Unavailable)),
_ => Err(Error::Unavailable),
}
}
}
type ScanPage = (u64, Vec<String>);
fn primary_pages(value: redis::Value) -> Result<Vec<(NodeAddress, ScanPage)>, Error> {
let redis::Value::Map(entries) = value else {
return Err(Error::Unavailable);
};
entries
.into_iter()
.map(|(node, page)| {
let node = redis::from_redis_value::<String>(node).map_err(|_| Error::Unavailable)?;
let node = NodeAddress::try_from(node.as_str()).map_err(|_| Error::Unavailable)?;
let page = redis::from_redis_value::<ScanPage>(page).map_err(|_| Error::Unavailable)?;
Ok((node, page))
})
.collect()
}
fn scan_command(cursor: u64, pattern: &str, count: usize) -> redis::Cmd {
let mut command = redis::cmd("SCAN");
command
.cursor_arg(cursor)
.arg("MATCH")
.arg(pattern)
.arg("COUNT")
.arg(count);
command
}
fn slot_groups(commands: &[redis::Cmd]) -> HashMap<Slot, Vec<usize>> {
let mut groups: HashMap<Slot, Vec<usize>> = HashMap::new();
for (index, command) in commands.iter().enumerate() {
let key = match command.args_iter().nth(1) {
Some(redis::Arg::Simple(key)) => key,
_ => b"",
};
groups.entry(Slot::for_key(key)).or_default().push(index);
}
groups
}

View file

@ -0,0 +1,632 @@
use std::{sync::Arc, time::Duration};
use litellm_cache::{
CacheCodec, CacheScript, ClientInfoCache, Error, IncrementOperation, QueueCache, ScanCache,
ScriptCache, SetCache, TtlCache,
};
use redis::Commands;
use super::{ConnectionRef, Connections, RedisCache, namespaced_key};
const INCREMENT_WITH_FLOOR_SCRIPT: &str = concat!(
"local count = redis.call('INCRBY', KEYS[1], ARGV[1]); ",
"if count < 0 then count = redis.call('INCRBY', KEYS[1], -count); end; ",
"if redis.call('TTL', KEYS[1]) < 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]); end; ",
"return count"
);
const SET_MAX_SCRIPT: &str = concat!(
"local current = redis.call('GET', KEYS[1]); ",
"if current == false or tonumber(current) < tonumber(ARGV[1]) then ",
"redis.call('SET', KEYS[1], ARGV[1]); ",
"if tonumber(ARGV[2]) > 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]); end; ",
"return ARGV[1]; end; return current"
);
#[derive(Clone, Debug, PartialEq)]
pub enum RedisArg {
Bytes(Vec<u8>),
Integer(i64),
Float(f64),
}
impl From<&str> for RedisArg {
fn from(value: &str) -> Self {
Self::Bytes(value.as_bytes().to_vec())
}
}
impl From<String> for RedisArg {
fn from(value: String) -> Self {
Self::Bytes(value.into_bytes())
}
}
impl From<Vec<u8>> for RedisArg {
fn from(value: Vec<u8>) -> Self {
Self::Bytes(value)
}
}
impl From<i64> for RedisArg {
fn from(value: i64) -> Self {
Self::Integer(value)
}
}
impl From<f64> for RedisArg {
fn from(value: f64) -> Self {
Self::Float(value)
}
}
impl redis::ToRedisArgs for RedisArg {
fn write_redis_args<W>(&self, out: &mut W)
where
W: ?Sized + redis::RedisWrite,
{
match self {
Self::Bytes(value) => value.write_redis_args(out),
Self::Integer(value) => value.write_redis_args(out),
Self::Float(value) => value.write_redis_args(out),
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct RedisRpushOperation {
pub key: String,
pub values: Vec<RedisArg>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RedisLpopOperation {
pub key: String,
pub count: Option<usize>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum RedisLpopResult {
Missing,
Value(Vec<u8>),
Values(Vec<Vec<u8>>),
}
pub struct RedisScript<C> {
connections: Arc<Connections<C>>,
namespace: Option<String>,
source: String,
}
impl<C> CacheScript for RedisScript<C>
where
C: redis::ConnectionLike + Send + 'static,
{
type Argument = RedisArg;
type Output = redis::Value;
async fn invoke(
&self,
keys: Vec<String>,
arguments: Vec<Self::Argument>,
) -> Result<Self::Output, Error> {
let keys = keys
.into_iter()
.map(|key| namespaced_key(self.namespace.as_deref(), &key))
.collect::<Vec<_>>();
let connections = Arc::clone(&self.connections);
let source = self.source.clone();
tokio::task::spawn_blocking(move || {
connections.execute(|connection| {
redis::cmd("EVAL")
.arg(source)
.arg(keys.len())
.arg(keys)
.arg(arguments)
.query(connection)
.map_err(|_| Error::Unavailable)
})
})
.await
.map_err(|_| Error::Unavailable)?
}
}
impl<S, C> RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
pub async fn delete_cache_keys(&self, keys: Vec<String>) -> Result<usize, Error> {
if keys.is_empty() {
return Ok(0);
}
let keys = keys
.into_iter()
.map(|key| self.namespaced_key(&key))
.collect::<Vec<_>>();
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
connection.del(keys).map_err(|_| Error::Unavailable)
})
.await
}
pub fn batch_get_counts(&self, keys: &[String]) -> Result<Vec<Option<i64>>, Error> {
let keys = keys
.iter()
.map(|key| self.namespaced_key(key))
.collect::<Vec<_>>();
let values = self.connections.execute(|connection| {
redis::cmd("MGET")
.arg(keys)
.query::<Vec<redis::Value>>(connection)
.map_err(|_| Error::Unavailable)
})?;
values.into_iter().map(count).collect()
}
pub async fn async_batch_get_counts(
&self,
keys: Vec<String>,
) -> Result<Vec<Option<i64>>, Error> {
let keys = keys
.iter()
.map(|key| self.namespaced_key(key))
.collect::<Vec<_>>();
let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("MGET")
.arg(keys)
.query::<Vec<redis::Value>>(connection)
.map_err(|_| Error::Unavailable)
})
.await?;
values.into_iter().map(count).collect()
}
pub fn sync_ping(&self) -> Result<bool, Error> {
self.connections
.execute(|connection| connection.ping().map_err(|_| Error::Unavailable))
}
pub async fn ping(&self) -> Result<bool, Error> {
Self::run_blocking(Arc::clone(&self.connections), |connection| {
connection.ping().map_err(|_| Error::Unavailable)
})
.await
}
pub async fn async_get_ttl(&self, key: &str) -> Result<Option<i64>, Error> {
let key = self.namespaced_key(key);
let ttl = Self::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("TTL")
.arg(key)
.query::<i64>(connection)
.map_err(|_| Error::Unavailable)
})
.await?;
Ok((ttl >= 0).then_some(ttl))
}
pub async fn async_scan_iter(&self, pattern: &str, count: usize) -> Result<Vec<String>, Error> {
let pattern = format!("{}*", self.namespaced_key(pattern));
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut matches = Vec::new();
connection.scan(&pattern, count, |_, keys| {
matches.extend(keys);
Ok(matches.len() < count)
})?;
matches.truncate(count);
Ok(matches)
})
.await
}
pub async fn async_set_cache_sadd(
&self,
key: &str,
values: Vec<RedisArg>,
ttl: Option<Duration>,
) -> Result<usize, Error> {
if values.is_empty() {
return Err(Error::InvalidEntry);
}
let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(ttl.unwrap_or(self.default_ttl));
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut sadd = redis::cmd("SADD");
sadd.arg(&key).arg(values);
let mut expire = redis::cmd("EXPIRE");
expire.arg(&key).arg(ttl);
let replies = connection.pipeline(vec![sadd, expire])?;
replies
.into_iter()
.next()
.map(redis::from_redis_value::<usize>)
.transpose()
.map_err(|_| Error::Unavailable)?
.ok_or(Error::Unavailable)
})
.await
}
pub async fn async_rpush(&self, key: &str, values: Vec<RedisArg>) -> Result<usize, Error> {
if values.is_empty() {
return Err(Error::InvalidEntry);
}
let key = self.namespaced_key(key);
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("RPUSH")
.arg(key)
.arg(values)
.query(connection)
.map_err(|_| Error::Unavailable)
})
.await
}
pub async fn async_rpush_pipeline(
&self,
operations: Vec<RedisRpushOperation>,
) -> Result<Vec<usize>, Error> {
let operations = operations
.into_iter()
.map(|operation| {
if operation.values.is_empty() {
return Err(Error::InvalidEntry);
}
Ok((self.namespaced_key(&operation.key), operation.values))
})
.collect::<Result<Vec<_>, _>>()?;
if operations.is_empty() {
return Ok(Vec::new());
}
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
let commands = operations
.into_iter()
.map(|(key, values)| {
let mut command = redis::cmd("RPUSH");
command.arg(key).arg(values);
command
})
.collect();
connection
.pipeline(commands)?
.into_iter()
.map(|value| redis::from_redis_value(value).map_err(|_| Error::Unavailable))
.collect()
})
.await
}
pub async fn async_lpop(
&self,
key: &str,
count: Option<usize>,
) -> Result<RedisLpopResult, Error> {
let key = self.namespaced_key(key);
let multiple = count.is_some();
let value = Self::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut command = redis::cmd("LPOP");
command.arg(key);
if let Some(count) = count {
command.arg(count);
}
command
.query::<redis::Value>(connection)
.map_err(|_| Error::Unavailable)
})
.await?;
lpop_result(value, multiple)
}
pub async fn async_lpop_pipeline(
&self,
operations: Vec<RedisLpopOperation>,
) -> Result<Vec<RedisLpopResult>, Error> {
let operations = operations
.into_iter()
.map(|operation| (self.namespaced_key(&operation.key), operation.count))
.collect::<Vec<_>>();
if operations.is_empty() {
return Ok(Vec::new());
}
let multiple = operations
.iter()
.map(|(_, count)| count.is_some())
.collect::<Vec<_>>();
let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| {
let commands = operations
.into_iter()
.map(|(key, count)| {
let mut command = redis::cmd("LPOP");
command.arg(key);
if let Some(count) = count {
command.arg(count);
}
command
})
.collect();
connection.pipeline(commands)
})
.await?;
values
.into_iter()
.zip(multiple)
.map(|(value, multiple)| lpop_result(value, multiple))
.collect()
}
pub async fn async_eval(
&self,
script: String,
keys: Vec<String>,
arguments: Vec<RedisArg>,
) -> Result<redis::Value, Error> {
let keys = keys
.into_iter()
.map(|key| self.namespaced_key(&key))
.collect::<Vec<_>>();
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("EVAL")
.arg(script)
.arg(keys.len())
.arg(keys)
.arg(arguments)
.query(connection)
.map_err(|_| Error::Unavailable)
})
.await
}
pub fn client_list(&self) -> Result<String, Error> {
self.connections
.execute(|connection| connection.node_text(redis::cmd("CLIENT").arg("LIST")))
}
pub fn info(&self) -> Result<String, Error> {
self.connections
.execute(|connection| connection.node_text(&redis::cmd("INFO")))
}
pub fn flushall(&self) -> Result<(), Error> {
self.connections.execute(|connection| connection.flushall())
}
}
impl<S, C> RedisCache<S, C>
where
S: CacheCodec<Value = f64>,
C: redis::ConnectionLike + Send + 'static,
{
pub fn increment_with_floor(
&self,
key: &str,
amount: i64,
ttl: Duration,
) -> Result<i64, Error> {
let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(ttl);
self.connections
.execute(|connection| increment_with_floor(connection, key, amount, ttl))
}
pub async fn async_increment_pipeline(
&self,
operations: Vec<IncrementOperation>,
) -> Result<Vec<f64>, Error> {
let operations = operations
.into_iter()
.map(|operation| {
(
self.namespaced_key(&operation.key),
operation.amount,
operation.ttl.map(Self::ttl_seconds),
)
})
.collect::<Vec<_>>();
if operations.is_empty() {
return Ok(Vec::new());
}
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut commands = Vec::with_capacity(operations.len() * 2);
let mut increments = Vec::with_capacity(operations.len());
for (key, amount, ttl) in operations {
let mut increment = redis::cmd("INCRBYFLOAT");
increment.arg(&key).arg(amount);
increments.push(commands.len());
commands.push(increment);
if let Some(ttl) = ttl {
let mut expire = redis::cmd("EXPIRE");
expire.arg(key).arg(ttl);
commands.push(expire);
}
}
let mut replies = connection.pipeline(commands)?;
increments
.into_iter()
.map(|index| {
redis::from_redis_value(std::mem::take(&mut replies[index]))
.map_err(|_| Error::Unavailable)
})
.collect()
})
.await
}
pub async fn async_increment_with_floor(
&self,
key: &str,
amount: i64,
ttl: Duration,
) -> Result<i64, Error> {
let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(ttl);
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
increment_with_floor(connection, key, amount, ttl)
})
.await
}
pub async fn async_set_max(
&self,
key: &str,
value: f64,
ttl: Option<Duration>,
) -> Result<f64, Error> {
let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(ttl.unwrap_or(self.default_ttl));
Self::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("EVAL")
.arg(SET_MAX_SCRIPT)
.arg(1)
.arg(key)
.arg(value)
.arg(ttl)
.query(connection)
.map_err(|_| Error::Unavailable)
})
.await
}
}
fn redis_bytes(value: redis::Value) -> Result<Vec<u8>, Error> {
match value {
redis::Value::BulkString(bytes) => Ok(bytes),
redis::Value::SimpleString(text) => Ok(text.into_bytes()),
_ => Err(Error::InvalidEntry),
}
}
fn lpop_result(value: redis::Value, multiple: bool) -> Result<RedisLpopResult, Error> {
match value {
redis::Value::Nil => Ok(RedisLpopResult::Missing),
redis::Value::Array(values) if multiple => values
.into_iter()
.map(redis_bytes)
.collect::<Result<Vec<_>, _>>()
.map(RedisLpopResult::Values),
value if !multiple => redis_bytes(value).map(RedisLpopResult::Value),
_ => Err(Error::InvalidEntry),
}
}
fn count(value: redis::Value) -> Result<Option<i64>, Error> {
match value {
redis::Value::Nil => Ok(None),
redis::Value::Int(value) => Ok(Some(value)),
redis::Value::BulkString(value) => std::str::from_utf8(&value)
.ok()
.and_then(|value| value.parse().ok())
.map(Some)
.ok_or(Error::InvalidEntry),
redis::Value::SimpleString(value) => {
value.parse().map(Some).map_err(|_| Error::InvalidEntry)
}
_ => Err(Error::InvalidEntry),
}
}
fn increment_with_floor(
connection: &mut ConnectionRef<'_>,
key: String,
amount: i64,
ttl: u64,
) -> Result<i64, Error> {
redis::cmd("EVAL")
.arg(INCREMENT_WITH_FLOOR_SCRIPT)
.arg(1)
.arg(key)
.arg(amount)
.arg(ttl)
.query(connection)
.map_err(|_| Error::Unavailable)
}
impl<S, C> TtlCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
async fn async_get_ttl(&self, key: &str) -> Result<Option<Duration>, Error> {
RedisCache::async_get_ttl(self, key)
.await
.map(|ttl| ttl.map(|seconds| Duration::from_secs(seconds as u64)))
}
}
impl<S, C> ScanCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
async fn async_scan_iter(&self, pattern: &str, count: usize) -> Result<Vec<String>, Error> {
RedisCache::async_scan_iter(self, pattern, count).await
}
}
impl<S, C> ClientInfoCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
type ClientList = String;
type Info = String;
fn client_list(&self) -> Result<Self::ClientList, Error> {
RedisCache::client_list(self)
}
fn info(&self) -> Result<Self::Info, Error> {
RedisCache::info(self)
}
}
impl<S, C> SetCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
type SetValue = RedisArg;
type SetResult = usize;
async fn async_set_cache_sadd(
&self,
key: &str,
values: Vec<Self::SetValue>,
ttl: Option<Duration>,
) -> Result<Self::SetResult, Error> {
RedisCache::async_set_cache_sadd(self, key, values, ttl).await
}
}
impl<S, C> QueueCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
type QueueValue = RedisArg;
type PopResult = RedisLpopResult;
async fn async_rpush(&self, key: &str, values: Vec<Self::QueueValue>) -> Result<usize, Error> {
RedisCache::async_rpush(self, key, values).await
}
async fn async_lpop(&self, key: &str, count: Option<usize>) -> Result<Self::PopResult, Error> {
RedisCache::async_lpop(self, key, count).await
}
}
impl<S, C> ScriptCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
type Script = RedisScript<C>;
fn async_register_script(&self, source: String) -> Self::Script {
RedisScript {
connections: Arc::clone(&self.connections),
namespace: self.namespace.clone(),
source,
}
}
}

View file

@ -1,3 +1,7 @@
mod cache;
mod topology;
pub use cache::RedisCache;
pub use cache::{
RedisArg, RedisCache, RedisLpopOperation, RedisLpopResult, RedisRpushOperation, RedisScript,
};
pub use topology::{RedisNode, RedisTopology};

View file

@ -0,0 +1,14 @@
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RedisNode {
pub host: String,
pub port: u16,
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub enum RedisTopology {
#[default]
Standalone,
Cluster {
startup_nodes: Vec<RedisNode>,
},
}

View file

@ -1,6 +1,703 @@
use litellm_cache_redis::RedisCache;
use std::time::Duration;
use litellm_cache::{
BaseCache, BatchCache, BatchEntry, CacheCodec, CacheConnectionStatus, CacheScript, ClaimCache,
CounterCache, DeleteCache, Error, ExactCacheContext, FlushCache, IncrementOperation, JsonCodec,
ScriptCache, get_cache, set_cache,
};
use litellm_cache_redis::{
RedisArg, RedisCache, RedisLpopOperation, RedisLpopResult, RedisRpushOperation,
};
use redis_test::{MockCmd, MockRedisConnection};
struct TaggedByteCodec(u8);
impl CacheCodec for TaggedByteCodec {
type Value = u8;
fn encode(&self, value: &u8) -> Result<Vec<u8>, Error> {
if *value > 127 {
return Err(Error::InvalidEntry);
}
Ok(vec![self.0, *value])
}
fn decode(&self, bytes: &[u8]) -> Result<u8, Error> {
match bytes {
[tag, value] if *tag == self.0 => Ok(*value),
_ => Err(Error::InvalidEntry),
}
}
}
#[test]
fn constructor_rejects_invalid_urls() {
assert!(RedisCache::new("not a redis url", None).is_err());
assert!(RedisCache::new("not a redis url", None, JsonCodec::<String>::new()).is_err());
}
#[test]
fn generic_helpers_use_the_injected_codec_and_ttl() {
let connection = MockRedisConnection::new([
MockCmd::new(
redis::cmd("SETEX")
.arg("counter")
.arg(2)
.arg([42u8, 7].as_slice()),
Ok("OK"),
),
MockCmd::new(redis::cmd("GET").arg("counter"), Ok(vec![42u8, 7])),
])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None, TaggedByteCodec(42));
let context = ExactCacheContext {
ttl: Some(Duration::from_millis(1500)),
};
set_cache(&cache, "counter", 7, &context).unwrap();
assert_eq!(get_cache(&cache, "counter", &context).unwrap(), Some(7));
}
#[tokio::test]
async fn async_operations_preserve_codec_ttl_and_missing_values() {
let connection = MockRedisConnection::new([
MockCmd::new(
redis::cmd("SETEX")
.arg("counter")
.arg(9)
.arg([42u8, 7].as_slice()),
Ok("OK"),
),
MockCmd::new(redis::cmd("GET").arg("counter"), Ok(vec![42u8, 7])),
MockCmd::new(
redis::cmd("SETEX")
.arg("batch")
.arg(2)
.arg([42u8, 8].as_slice()),
Ok("OK"),
),
MockCmd::new(redis::cmd("DEL").arg("counter"), Ok(1u32)),
MockCmd::new(redis::cmd("GET").arg("counter"), Ok(redis::Value::Nil)),
])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(
connection,
Some(Duration::from_secs(9)),
TaggedByteCodec(42),
);
let context = ExactCacheContext::default();
cache
.batch_cache_write("counter", 7, context.clone())
.await
.unwrap();
assert_eq!(
cache.async_get_cache("counter", &context).await.unwrap(),
Some(7)
);
cache
.async_set_cache_pipeline(
vec![("batch".into(), 8)],
ExactCacheContext {
ttl: Some(Duration::from_millis(1500)),
},
)
.await
.unwrap();
cache.async_delete_cache("counter").await.unwrap();
assert_eq!(
cache.async_get_cache("counter", &context).await.unwrap(),
None
);
}
#[tokio::test]
async fn codec_errors_propagate_without_writing_partial_batches() {
let connection = MockRedisConnection::new([
MockCmd::new(redis::cmd("GET").arg("invalid"), Ok(vec![99u8, 7])),
MockCmd::new(redis::cmd("GET").arg("invalid"), Ok(vec![99u8, 7])),
])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None, TaggedByteCodec(42));
let context = ExactCacheContext::default();
assert_eq!(
cache.set_cache("invalid", 255, &context),
Err(Error::InvalidEntry)
);
assert_eq!(
cache.async_set_cache("invalid", 255, context.clone()).await,
Err(Error::InvalidEntry)
);
assert_eq!(
cache
.async_set_cache_pipeline(
vec![("valid".into(), 7), ("invalid".into(), 255)],
context.clone(),
)
.await,
Err(Error::InvalidEntry)
);
assert_eq!(
cache.get_cache("invalid", &context),
Err(Error::InvalidEntry)
);
assert_eq!(
cache.async_get_cache("invalid", &context).await,
Err(Error::InvalidEntry)
);
}
#[test]
fn namespaces_are_optional_and_existing_prefixes_are_not_duplicated() {
let connection = MockRedisConnection::new([
MockCmd::new(redis::cmd("GET").arg("team:key"), Ok(redis::Value::Nil)),
MockCmd::new(redis::cmd("GET").arg("team:key"), Ok(redis::Value::Nil)),
])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None, JsonCodec::<String>::new())
.with_namespace(Some("team".into()));
assert_eq!(
cache
.get_cache("key", &ExactCacheContext::default())
.unwrap(),
None
);
assert_eq!(
cache
.get_cache("team:key", &ExactCacheContext::default())
.unwrap(),
None
);
}
#[test]
fn flush_requires_a_namespace_and_escapes_glob_metacharacters() {
let unscoped = RedisCache::with_connection(
MockRedisConnection::new([]).assert_all_commands_consumed(),
None,
JsonCodec::<String>::new(),
);
assert_eq!(unscoped.flush_cache(), Err(Error::UnscopedFlush));
let connection = MockRedisConnection::new([
MockCmd::new(
redis::cmd("SCAN")
.cursor_arg(0)
.arg("MATCH")
.arg("team\\*:*")
.arg("COUNT")
.arg(1000),
Ok(redis_test::redis_value!(["0", ["team*:key"]])),
),
MockCmd::new(redis::cmd("DEL").arg("team*:key"), Ok(1u32)),
])
.assert_all_commands_consumed();
let scoped = RedisCache::with_connection(connection, None, JsonCodec::<String>::new())
.with_namespace(Some("team*".into()));
scoped.flush_cache().unwrap();
}
#[tokio::test]
async fn connection_failures_use_the_python_result_contract() {
let error = redis::RedisError::from((redis::ErrorKind::Io, "connection refused"));
let connection =
MockRedisConnection::new([MockCmd::new(redis::cmd("PING"), Err::<String, _>(error))])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None, JsonCodec::<String>::new());
let result = cache.test_connection().await.unwrap();
assert_eq!(result.status, CacheConnectionStatus::Failed);
assert!(result.message.starts_with("Redis connection failed:"));
assert!(result.error.is_some());
}
#[tokio::test]
async fn batch_reads_keep_order_and_treat_invalid_values_as_invalid_entries() {
let connection = MockRedisConnection::new([MockCmd::new(
redis::cmd("MGET").arg("hit").arg("miss").arg("invalid"),
Ok(vec![
redis::Value::BulkString(vec![42, 7]),
redis::Value::Nil,
redis::Value::BulkString(vec![99, 7]),
]),
)])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None, TaggedByteCodec(42));
assert_eq!(
cache
.async_batch_get_cache(
vec!["hit".into(), "miss".into(), "invalid".into()],
ExactCacheContext::default(),
)
.await
.unwrap(),
vec![BatchEntry::Hit(7), BatchEntry::Miss, BatchEntry::Invalid]
);
}
#[tokio::test]
async fn async_flush_deletes_each_scan_page_separately() {
let connection = MockRedisConnection::new([
MockCmd::new(
redis::cmd("SCAN")
.cursor_arg(0)
.arg("MATCH")
.arg("team:*")
.arg("COUNT")
.arg(1000),
Ok(redis_test::redis_value!(["7", ["team:a", "team:b"]])),
),
MockCmd::new(redis::cmd("DEL").arg("team:a").arg("team:b"), Ok(2u32)),
MockCmd::new(
redis::cmd("SCAN")
.cursor_arg(7)
.arg("MATCH")
.arg("team:*")
.arg("COUNT")
.arg(1000),
Ok(redis_test::redis_value!(["0", ["team:c"]])),
),
MockCmd::new(redis::cmd("DEL").arg("team:c"), Ok(1u32)),
])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None, JsonCodec::<String>::new())
.with_namespace(Some("team".into()));
cache.async_flush_cache().await.unwrap();
}
#[tokio::test]
async fn direct_redis_operations_preserve_namespace_values_and_missing_ttls() {
let mut sadd_pipeline = redis::pipe();
sadd_pipeline
.cmd("SADD")
.arg("team:members")
.arg("a")
.arg("b")
.cmd("EXPIRE")
.arg("team:members")
.arg(600u64)
.ignore();
let connection = MockRedisConnection::new([
MockCmd::new(
redis::cmd("MGET").arg("team:count").arg("team:missing"),
Ok(redis_test::redis_value!(["7", nil])),
),
MockCmd::new(
redis::cmd("MGET").arg("team:count").arg("team:missing"),
Ok(redis_test::redis_value!(["7", nil])),
),
MockCmd::new(redis::cmd("PING"), Ok("PONG")),
MockCmd::new(redis::cmd("PING"), Ok("PONG")),
MockCmd::new(redis::cmd("TTL").arg("team:missing"), Ok(-2i64)),
MockCmd::new(
redis::cmd("SCAN")
.cursor_arg(0)
.arg("MATCH")
.arg("team:job-*")
.arg("COUNT")
.arg(25),
Ok(redis_test::redis_value!(["4", ["team:job-a"]])),
),
MockCmd::new(
redis::cmd("SCAN")
.cursor_arg(4)
.arg("MATCH")
.arg("team:job-*")
.arg("COUNT")
.arg(25),
Ok(redis_test::redis_value!(["0", ["team:job-b"]])),
),
MockCmd::new(
redis::cmd("DEL").arg("team:job-a").arg("team:job-b"),
Ok(2u32),
),
MockCmd::with_values(
sadd_pipeline,
Ok(vec![redis::Value::Int(2), redis::Value::Int(1)]),
),
MockCmd::new(
redis::cmd("RPUSH").arg("team:queue").arg("a").arg("b"),
Ok(2u32),
),
MockCmd::new(
redis::cmd("LPOP").arg("team:queue").arg(2usize),
Ok(redis_test::redis_value!(["a", "b"])),
),
MockCmd::new(
redis::cmd("EVAL")
.arg("return KEYS[1]")
.arg(1usize)
.arg("team:key"),
Ok("team:key"),
),
MockCmd::new(
redis::cmd("EVAL")
.arg("return KEYS[1]")
.arg(1usize)
.arg("team:key"),
Ok("team:key"),
),
MockCmd::new(redis::cmd("CLIENT").arg("LIST"), Ok("id=1")),
MockCmd::new(redis::cmd("INFO"), Ok("redis_version:7")),
MockCmd::new(redis::cmd("FLUSHALL"), Ok("OK")),
])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None, JsonCodec::<String>::new())
.with_namespace(Some("team".into()));
assert_eq!(
cache
.batch_get_counts(&["count".into(), "missing".into()])
.unwrap(),
[Some(7), None]
);
assert_eq!(
cache
.async_batch_get_counts(vec!["count".into(), "missing".into()])
.await
.unwrap(),
[Some(7), None]
);
assert!(cache.sync_ping().unwrap());
assert!(cache.ping().await.unwrap());
assert_eq!(cache.async_get_ttl("missing").await.unwrap(), None);
assert_eq!(
cache.async_scan_iter("job-", 25).await.unwrap(),
["team:job-a", "team:job-b"]
);
assert_eq!(
cache
.delete_cache_keys(vec!["job-a".into(), "job-b".into()])
.await
.unwrap(),
2
);
assert_eq!(
cache
.async_set_cache_sadd("members", vec!["a".into(), "b".into()], None)
.await
.unwrap(),
2
);
assert_eq!(
cache
.async_rpush("queue", vec!["a".into(), "b".into()])
.await
.unwrap(),
2
);
assert_eq!(
cache.async_lpop("queue", Some(2)).await.unwrap(),
RedisLpopResult::Values(vec![b"a".to_vec(), b"b".to_vec()])
);
assert_eq!(
cache
.async_eval("return KEYS[1]".into(), vec!["key".into()], Vec::new())
.await
.unwrap(),
redis::Value::BulkString(b"team:key".to_vec())
);
assert_eq!(
cache
.async_register_script("return KEYS[1]".into())
.invoke(vec!["key".into()], Vec::new())
.await
.unwrap(),
redis::Value::BulkString(b"team:key".to_vec())
);
assert_eq!(cache.client_list().unwrap(), "id=1");
assert_eq!(cache.info().unwrap(), "redis_version:7");
cache.flushall().unwrap();
}
#[tokio::test]
async fn direct_redis_pipelines_preserve_operation_order() {
let mut rpush_pipeline = redis::pipe();
rpush_pipeline
.cmd("RPUSH")
.arg("team:a")
.arg("one")
.cmd("RPUSH")
.arg("team:b")
.arg("two");
let mut lpop_pipeline = redis::pipe();
lpop_pipeline
.cmd("LPOP")
.arg("team:a")
.arg(2usize)
.cmd("LPOP")
.arg("team:b");
let connection = MockRedisConnection::new([
MockCmd::with_values(
rpush_pipeline,
Ok(vec![redis::Value::Int(1), redis::Value::Int(2)]),
),
MockCmd::with_values(
lpop_pipeline,
Ok(vec![redis_test::redis_value!(["one"]), redis::Value::Nil]),
),
])
.assert_all_commands_consumed();
let queue = RedisCache::with_connection(connection, None, JsonCodec::<String>::new())
.with_namespace(Some("team".into()));
assert_eq!(
queue
.async_rpush_pipeline(vec![
RedisRpushOperation {
key: "a".into(),
values: vec![RedisArg::from("one")],
},
RedisRpushOperation {
key: "b".into(),
values: vec![RedisArg::from("two")],
},
])
.await
.unwrap(),
[1, 2]
);
assert_eq!(
queue
.async_lpop_pipeline(vec![
RedisLpopOperation {
key: "a".into(),
count: Some(2),
},
RedisLpopOperation {
key: "b".into(),
count: None,
},
])
.await
.unwrap(),
[
RedisLpopResult::Values(vec![b"one".to_vec()]),
RedisLpopResult::Missing,
]
);
let mut increment_pipeline = redis::pipe();
increment_pipeline
.cmd("INCRBYFLOAT")
.arg("team:counter")
.arg(1.5f64)
.cmd("EXPIRE")
.arg("team:counter")
.arg(10u64)
.ignore()
.cmd("INCRBYFLOAT")
.arg("team:counter")
.arg(2.0f64);
let connection = MockRedisConnection::new([MockCmd::with_values(
increment_pipeline,
Ok(vec![
redis::Value::BulkString(b"1.5".to_vec()),
redis::Value::Int(1),
redis::Value::BulkString(b"3.5".to_vec()),
]),
)])
.assert_all_commands_consumed();
let counters = RedisCache::with_connection(connection, None, JsonCodec::<f64>::new())
.with_namespace(Some("team".into()));
assert_eq!(
counters
.async_increment_pipeline(vec![
IncrementOperation {
key: "counter".into(),
amount: 1.5,
ttl: Some(Duration::from_secs(10)),
},
IncrementOperation {
key: "counter".into(),
amount: 2.0,
ttl: None,
},
])
.await
.unwrap(),
[1.5, 3.5]
);
}
const INCREMENT_WITH_FLOOR_SCRIPT: &str = concat!(
"local count = redis.call('INCRBY', KEYS[1], ARGV[1]); ",
"if count < 0 then count = redis.call('INCRBY', KEYS[1], -count); end; ",
"if redis.call('TTL', KEYS[1]) < 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]); end; ",
"return count"
);
const SET_MAX_SCRIPT: &str = concat!(
"local current = redis.call('GET', KEYS[1]); ",
"if current == false or tonumber(current) < tonumber(ARGV[1]) then ",
"redis.call('SET', KEYS[1], ARGV[1]); ",
"if tonumber(ARGV[2]) > 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]); end; ",
"return ARGV[1]; end; return current"
);
#[tokio::test]
async fn counter_repairs_are_atomic_and_use_default_ttl() {
let floor = || {
redis::cmd("EVAL")
.arg(INCREMENT_WITH_FLOOR_SCRIPT)
.arg(1)
.arg("team:counter")
.arg(-2i64)
.arg(30u64)
.clone()
};
let connection = MockRedisConnection::new([
MockCmd::new(floor(), Ok(0i64)),
MockCmd::new(floor(), Ok(0i64)),
MockCmd::new(
redis::cmd("EVAL")
.arg(SET_MAX_SCRIPT)
.arg(1)
.arg("team:counter")
.arg(4.5f64)
.arg(600u64),
Ok("4.5"),
),
])
.assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None, JsonCodec::<f64>::new())
.with_namespace(Some("team".into()));
assert_eq!(
cache
.increment_with_floor("counter", -2, Duration::from_secs(30))
.unwrap(),
0
);
assert_eq!(
cache
.async_increment_with_floor("counter", -2, Duration::from_secs(30))
.await
.unwrap(),
0
);
assert_eq!(
cache.async_set_max("counter", 4.5, None).await.unwrap(),
4.5
);
}
const CLAIM_SCRIPT: &str = concat!(
"local current = redis.call('GET', KEYS[1]); ",
"if ARGV[1] == '' then if current ~= false and current ~= '' then return 0; end; ",
"elseif current ~= ARGV[1] then return 0; end; ",
"if ARGV[3] ~= '' then redis.call('SET', KEYS[1], ARGV[3], 'EX', ARGV[2]); ",
"elseif ARGV[4] == '1' then redis.call('EXPIRE', KEYS[1], ARGV[2]); end; return 1"
);
fn claim_eval(expected: &str, write: &str, refresh: bool) -> redis::Cmd {
let mut cmd = redis::cmd("EVAL");
cmd.arg(CLAIM_SCRIPT)
.arg(1)
.arg("pin")
.arg(expected)
.arg(600)
.arg(write)
.arg(u8::from(refresh));
cmd
}
#[tokio::test]
async fn claims_match_eligible_values_written_by_another_encoder() {
let python_payload = r#"{"model_id": "a", "deployment": "east"}"#;
let stored = serde_json::json!({"deployment": "east", "model_id": "a"});
let candidate = serde_json::json!({"model_id": "b"});
let connection = MockRedisConnection::new([
MockCmd::new(redis::cmd("GET").arg("pin"), Ok(python_payload)),
MockCmd::new(claim_eval(python_payload, "", true), Ok(1)),
])
.assert_all_commands_consumed();
let cache =
RedisCache::with_connection(connection, None, JsonCodec::<serde_json::Value>::new());
assert_eq!(
cache
.async_claim_cache(
"pin",
candidate,
vec![stored.clone()],
ExactCacheContext::default()
)
.await
.unwrap(),
stored
);
}
#[test]
fn claims_retry_when_the_key_changes_and_replace_ineligible_winners() {
let candidate = serde_json::json!({"model_id": "b"});
let payload = r#"{"model_id":"b"}"#;
let connection = MockRedisConnection::new([
MockCmd::new(redis::cmd("GET").arg("pin"), Ok(redis::Value::Nil)),
MockCmd::new(claim_eval("", payload, false), Ok(0)),
MockCmd::new(redis::cmd("GET").arg("pin"), Ok(r#"{"model_id":"gone"}"#)),
MockCmd::new(claim_eval(r#"{"model_id":"gone"}"#, payload, false), Ok(1)),
])
.assert_all_commands_consumed();
let cache =
RedisCache::with_connection(connection, None, JsonCodec::<serde_json::Value>::new());
assert_eq!(
cache
.claim_cache(
"pin",
candidate.clone(),
&[serde_json::json!({"model_id": "a"})],
ExactCacheContext::default()
)
.unwrap(),
candidate
);
}
#[test]
fn claims_without_eligible_values_keep_the_winner_without_refreshing_its_ttl() {
let stored = r#"{"model_id": "a"}"#;
let connection = MockRedisConnection::new([
MockCmd::new(redis::cmd("GET").arg("pin"), Ok(stored)),
MockCmd::new(claim_eval(stored, "", false), Ok(1)),
])
.assert_all_commands_consumed();
let cache =
RedisCache::with_connection(connection, None, JsonCodec::<serde_json::Value>::new());
assert_eq!(
cache
.claim_cache(
"pin",
serde_json::json!({"model_id": "b"}),
&[],
ExactCacheContext::default()
)
.unwrap(),
serde_json::json!({"model_id": "a"})
);
}
#[tokio::test]
async fn async_increment_runs_the_atomic_script() {
let mut eval = redis::cmd("EVAL");
eval.arg(concat!(
"local value = redis.call('INCRBYFLOAT', KEYS[1], ARGV[1]); ",
"if redis.call('TTL', KEYS[1]) == -1 then ",
"redis.call('EXPIRE', KEYS[1], ARGV[2]); end; return value"
))
.arg(1)
.arg("counter")
.arg(2.5f64)
.arg(600);
let connection =
MockRedisConnection::new([MockCmd::new(eval, Ok("4.5"))]).assert_all_commands_consumed();
let cache = RedisCache::with_connection(connection, None, JsonCodec::<f64>::new());
assert_eq!(
cache
.async_increment("counter", 2.5, ExactCacheContext::default())
.await
.unwrap(),
4.5
);
}

View file

@ -0,0 +1,492 @@
//! Contract tests against a real Redis Cluster. Set `LITELLM_TEST_REDIS_CLUSTER_NODES` to a
//! comma separated `host:port` list (for example `127.0.0.1:7000,127.0.0.1:7001`) to run them.
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use litellm_cache::{
BaseCache, BatchCache, BatchEntry, CacheConnectionStatus, CacheScript, ClaimCache,
CounterCache, DeleteCache, Error, ExactCacheContext, FlushCache, IncrementOperation, JsonCodec,
ScriptCache,
};
use litellm_cache_redis::{
RedisArg, RedisCache, RedisLpopOperation, RedisLpopResult, RedisNode, RedisRpushOperation,
RedisTopology,
};
use redis::cluster_routing::Slot;
type Cache = RedisCache<JsonCodec<serde_json::Value>>;
fn topology() -> Option<RedisTopology> {
let nodes = std::env::var("LITELLM_TEST_REDIS_CLUSTER_NODES").ok()?;
let startup_nodes = nodes
.split(',')
.map(|node| {
let (host, port) = node.trim().rsplit_once(':').expect("host:port");
RedisNode {
host: host.to_string(),
port: port.parse().expect("port"),
}
})
.collect();
Some(RedisTopology::Cluster { startup_nodes })
}
fn namespace(label: &str) -> String {
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos();
format!("cluster-test:{label}:{nanos}")
}
fn cluster_url() -> String {
std::env::var("LITELLM_TEST_REDIS_CLUSTER_URL")
.unwrap_or_else(|_| "redis://127.0.0.1:7000".into())
}
fn cluster_cache(label: &str) -> Option<Cache> {
let topology = topology()?;
Some(
Cache::connect(
&cluster_url(),
&topology,
Some(Duration::from_secs(120)),
JsonCodec::new(),
)
.expect("cluster connection")
.with_namespace(Some(namespace(label))),
)
}
fn counter_cache(label: &str) -> Option<RedisCache<JsonCodec<f64>>> {
let topology = topology()?;
Some(
RedisCache::connect(
&cluster_url(),
&topology,
Some(Duration::from_secs(60)),
JsonCodec::new(),
)
.expect("cluster connection")
.with_namespace(Some(namespace(label))),
)
}
fn multi_slot_keys(count: usize) -> Vec<String> {
let keys: Vec<String> = (0..count).map(|index| format!("key-{index}")).collect();
let slots: std::collections::HashSet<Slot> = keys.iter().map(Slot::for_key).collect();
assert!(slots.len() > 1, "keys must span multiple slots");
keys
}
macro_rules! cluster_or_skip {
($label:expr) => {
match cluster_cache($label) {
Some(cache) => cache,
None => return,
}
};
}
#[test]
fn constructor_rejects_clusters_without_startup_nodes() {
let error = Cache::connect(
"redis://127.0.0.1:7000",
&RedisTopology::Cluster {
startup_nodes: Vec::new(),
},
None,
JsonCodec::new(),
)
.err();
assert!(matches!(error, Some(Error::Unavailable)));
}
#[test]
fn constructor_rejects_unix_socket_urls_for_clusters() {
let error = Cache::connect(
"redis+unix:///tmp/redis.sock",
&RedisTopology::Cluster {
startup_nodes: vec![RedisNode {
host: "127.0.0.1".into(),
port: 7000,
}],
},
None,
JsonCodec::new(),
)
.err();
assert!(matches!(error, Some(Error::Unavailable)));
}
#[test]
fn single_key_operations_round_trip_with_ttl_rounding() {
let cache = cluster_or_skip!("single");
let context = ExactCacheContext {
ttl: Some(Duration::from_millis(1500)),
};
let keys = multi_slot_keys(12);
for (index, key) in keys.iter().enumerate() {
cache
.set_cache(key, serde_json::json!({ "index": index }), &context)
.unwrap();
}
for (index, key) in keys.iter().enumerate() {
assert_eq!(
cache.get_cache(key, &context).unwrap(),
Some(serde_json::json!({ "index": index }))
);
}
let runtime = tokio::runtime::Runtime::new().unwrap();
let ttl = runtime.block_on(cache.async_get_ttl(&keys[0])).unwrap();
assert_eq!(ttl, Some(2));
cache.delete_cache(&keys[0]).unwrap();
assert_eq!(cache.get_cache(&keys[0], &context).unwrap(), None);
assert!(cache.sync_ping().unwrap());
}
#[tokio::test]
async fn batch_reads_span_slots_and_preserve_order_with_malformed_entries() {
let cache = cluster_or_skip!("batch");
let context = ExactCacheContext::default();
let keys = multi_slot_keys(40);
for (index, key) in keys.iter().enumerate() {
if index % 5 == 0 {
continue;
}
cache
.async_set_cache(key, serde_json::json!(index), context.clone())
.await
.unwrap();
}
let mut raw = redis::cluster::ClusterClient::new(vec![cluster_url()])
.unwrap()
.get_connection()
.unwrap();
let malformed = format!("{}:{}", cache.namespace().unwrap(), keys[1]);
redis::cmd("SET")
.arg(&malformed)
.arg("not json")
.exec(&mut raw)
.unwrap();
let entries = cache
.async_batch_get_cache(keys.clone(), context.clone())
.await
.unwrap();
assert_eq!(entries.len(), keys.len());
for (index, entry) in entries.iter().enumerate() {
let expected = if index == 1 {
BatchEntry::Invalid
} else if index % 5 == 0 {
BatchEntry::Miss
} else {
BatchEntry::Hit(serde_json::json!(index))
};
assert_eq!(*entry, expected, "entry {index}");
}
let sync_entries = cache.batch_get_cache(&keys, &context).unwrap();
assert_eq!(sync_entries, entries);
cache.delete_cache_keys(keys.clone()).await.unwrap();
let entries = cache.async_batch_get_cache(keys, context).await.unwrap();
assert!(entries.iter().all(|entry| *entry == BatchEntry::Miss));
}
#[tokio::test]
async fn pipelines_group_by_slot_and_return_results_in_submission_order() {
let cache = cluster_or_skip!("pipeline");
let keys = multi_slot_keys(30);
let entries = keys
.iter()
.enumerate()
.map(|(index, key)| (key.clone(), serde_json::json!(index)))
.collect();
cache
.async_set_cache_pipeline(entries, ExactCacheContext::default())
.await
.unwrap();
let hits = cache
.async_batch_get_cache(keys.clone(), ExactCacheContext::default())
.await
.unwrap();
assert!(
hits.iter()
.enumerate()
.all(|(index, entry)| *entry == BatchEntry::Hit(serde_json::json!(index)))
);
let queues: Vec<String> = keys.iter().map(|key| format!("queue:{key}")).collect();
let pushed = cache
.async_rpush_pipeline(
queues
.iter()
.enumerate()
.map(|(index, key)| RedisRpushOperation {
key: key.clone(),
values: (0..=index)
.map(|value| RedisArg::Integer(value as i64))
.collect(),
})
.collect(),
)
.await
.unwrap();
assert_eq!(pushed, (1..=keys.len()).collect::<Vec<_>>());
let popped = cache
.async_lpop_pipeline(
queues
.iter()
.enumerate()
.map(|(index, key)| RedisLpopOperation {
key: key.clone(),
count: (index % 2 == 0).then_some(2),
})
.collect(),
)
.await
.unwrap();
for (index, result) in popped.into_iter().enumerate() {
match result {
RedisLpopResult::Value(value) => {
assert_eq!(index % 2, 1, "queue {index}");
assert_eq!(value, b"0");
}
RedisLpopResult::Values(values) => {
assert_eq!(index % 2, 0, "queue {index}");
let expected: Vec<Vec<u8>> = (0..=index)
.take(2)
.map(|value| value.to_string().into_bytes())
.collect();
assert_eq!(values, expected);
}
other => panic!("queue {index}: {other:?}"),
}
}
let counters: Vec<String> = keys.iter().map(|key| format!("counter:{key}")).collect();
let Some(counter) = counter_cache("counter") else {
return;
};
let totals = counter
.async_increment_pipeline(
counters
.iter()
.enumerate()
.map(|(index, key)| IncrementOperation {
key: key.clone(),
amount: index as f64 + 0.5,
ttl: (index % 3 == 0).then_some(Duration::from_secs(30)),
})
.collect(),
)
.await
.unwrap();
let expected: Vec<f64> = (0..keys.len()).map(|index| index as f64 + 0.5).collect();
assert_eq!(totals, expected);
assert_eq!(counter.async_get_ttl(&counters[0]).await.unwrap(), Some(30));
assert_eq!(counter.async_get_ttl(&counters[1]).await.unwrap(), None);
counter.async_flush_cache().await.unwrap();
cache.async_flush_cache().await.unwrap();
}
#[tokio::test]
async fn scan_and_scoped_flush_cover_every_primary() {
let cache = cluster_or_skip!("flush");
let other = cluster_or_skip!("other");
let context = ExactCacheContext::default();
let keys = multi_slot_keys(60);
for key in &keys {
cache
.async_set_cache(key, serde_json::json!(true), context.clone())
.await
.unwrap();
other
.async_set_cache(key, serde_json::json!(true), context.clone())
.await
.unwrap();
}
let mut scanned = cache.async_scan_iter("key-", 1000).await.unwrap();
scanned.sort();
let mut expected: Vec<String> = keys
.iter()
.map(|key| format!("{}:{key}", cache.namespace().unwrap()))
.collect();
expected.sort();
assert_eq!(scanned, expected);
assert_eq!(cache.async_scan_iter("key-", 7).await.unwrap().len(), 7);
cache.flush_cache().unwrap();
let flushed = cache
.async_batch_get_cache(keys.clone(), context.clone())
.await
.unwrap();
assert!(flushed.iter().all(|entry| *entry == BatchEntry::Miss));
let kept = other.async_batch_get_cache(keys, context).await.unwrap();
assert!(
kept.iter()
.all(|entry| *entry == BatchEntry::Hit(serde_json::json!(true)))
);
other.async_flush_cache().await.unwrap();
}
fn ping_calls_per_node(startup: &redis::Client) -> Vec<(String, u64)> {
let mut connection = startup.get_connection().unwrap();
let nodes: String = redis::cmd("CLUSTER")
.arg("NODES")
.query(&mut connection)
.unwrap();
let mut counts: Vec<(String, u64)> = nodes
.lines()
.map(|line| {
let address = line.split_whitespace().nth(1).unwrap();
let address = address.split('@').next().unwrap();
let mut node = redis::Client::open(format!("redis://{address}"))
.unwrap()
.get_connection()
.unwrap();
let stats: String = redis::cmd("INFO")
.arg("commandstats")
.query(&mut node)
.unwrap();
let calls = stats
.lines()
.find_map(|stat| stat.strip_prefix("cmdstat_ping:calls="))
.and_then(|rest| rest.split(',').next())
.map_or(0, |calls| calls.parse().unwrap());
(address.to_string(), calls)
})
.collect();
counts.sort();
counts
}
#[tokio::test]
async fn ping_reaches_every_node() {
let cache = cluster_or_skip!("ping");
let startup = redis::Client::open(cluster_url()).unwrap();
let before = ping_calls_per_node(&startup);
assert!(before.len() >= 2, "{before:?}");
assert!(cache.ping().await.unwrap());
let after = ping_calls_per_node(&startup);
for ((node, calls_before), (_, calls_after)) in before.iter().zip(&after) {
assert!(calls_after > calls_before, "{node} was not pinged");
}
assert!(cache.sync_ping().unwrap());
let result = cache.test_connection().await.unwrap();
assert_eq!(result.status, CacheConnectionStatus::Success);
}
#[tokio::test]
async fn counters_claims_scripts_and_sets_work_on_the_cluster() {
let Some(counter) = counter_cache("counter") else {
return;
};
let context = ExactCacheContext::default();
assert_eq!(
counter
.increment_cache("spend", 1.5, context.clone())
.unwrap(),
1.5
);
assert_eq!(
counter
.async_increment("spend", 2.0, context.clone())
.await
.unwrap(),
3.5
);
assert_eq!(
counter
.increment_with_floor("budget", -3, Duration::from_secs(30))
.unwrap(),
0
);
assert_eq!(
counter
.async_increment_with_floor("budget", 7, Duration::from_secs(30))
.await
.unwrap(),
7
);
assert_eq!(counter.async_set_max("peak", 4.0, None).await.unwrap(), 4.0);
assert_eq!(counter.async_set_max("peak", 2.0, None).await.unwrap(), 4.0);
counter.flush_cache().unwrap();
let cache = cluster_or_skip!("claim");
let owner = serde_json::json!("owner-a");
let rival = serde_json::json!("owner-b");
assert_eq!(
cache
.claim_cache("lock", owner.clone(), &[], context.clone())
.unwrap(),
owner
);
assert_eq!(
cache
.async_claim_cache("lock", rival.clone(), vec![owner.clone()], context.clone())
.await
.unwrap(),
owner
);
assert_eq!(
cache
.claim_cache("lock", rival.clone(), &[], context.clone())
.unwrap(),
owner
);
assert_eq!(
cache
.async_claim_cache("lock", rival.clone(), vec![rival.clone()], context.clone())
.await
.unwrap(),
rival
);
let script = cache
.async_register_script("return redis.call('SET', KEYS[1], ARGV[1], 'EX', ARGV[2])".into());
let reply = script
.invoke(
vec!["scripted".into()],
vec![RedisArg::Bytes(b"payload".to_vec()), RedisArg::Integer(5)],
)
.await
.unwrap();
assert_eq!(reply, redis::Value::Okay);
assert_eq!(cache.async_get_ttl("scripted").await.unwrap(), Some(5));
let evaluated: redis::Value = cache
.async_eval(
"return redis.call('GET', KEYS[1])".into(),
vec!["scripted".into()],
Vec::new(),
)
.await
.unwrap();
assert_eq!(evaluated, redis::Value::BulkString(b"payload".to_vec()));
assert_eq!(
cache
.async_set_cache_sadd(
"members",
vec![
RedisArg::Bytes(b"a".to_vec()),
RedisArg::Bytes(b"b".to_vec())
],
Some(Duration::from_secs(9)),
)
.await
.unwrap(),
2
);
assert_eq!(cache.async_get_ttl("members").await.unwrap(), Some(9));
let result = cache.test_connection().await.unwrap();
assert_eq!(result.status, CacheConnectionStatus::Success);
assert!(cache.ping().await.unwrap());
let info = cache.info().unwrap();
assert!(info.matches("redis_version").count() > 1, "{info}");
assert!(cache.client_list().unwrap().contains("id="));
cache.async_flush_cache().await.unwrap();
assert_eq!(cache.async_get_ttl("members").await.unwrap(), None);
assert_eq!(cache.get_cache("lock", &context).unwrap(), None);
}

View file

@ -0,0 +1,20 @@
[package]
name = "litellm-cache-response"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-cache.workspace = true
py_literal = "0.4.0"
serde.workspace = true
serde_json.workspace = true
sha2.workspace = true
[dev-dependencies]
litellm-cache-memory.workspace = true
litellm-cache-redis.workspace = true
redis = "1.7.0"
redis-test = "1.0.4"
tokio.workspace = true

View file

@ -0,0 +1,61 @@
# Response cache foundation
`ResponseCache<B>` adds request keys, independent read/write controls, response envelopes, and freshness checks to any `B: BaseCache<Value = CacheEntry>`
## Ownership
`litellm-cache` defines typed storage and codec traits. Memory and Redis implement those traits without depending on response policy. Other consumers can store their own value types using the same backend implementations
`litellm-cache-response` owns response keys, controls, entries, the Python-compatible response codec, and `WriteBuffer`, the backend-neutral deferred-write policy. It has no runtime dependency on a specific cache backend or Python
The Python bridge constructs backends and selects them through its private `NativeResponseCache` enum, which only dispatches. Generic Rust callers inject their backend directly. A native gateway can construct the same generic response service in its own host
## Native Rust use
```rust
use std::{sync::Arc, time::Duration};
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{CacheKeyInput, ResponseCache, ResponseCacheRequest};
use serde_json::json;
let cache = ResponseCache::new(Arc::new(InMemoryCache::default()));
let request = ResponseCacheRequest::new(CacheKeyInput {
preset: Some("example:key".into()),
..Default::default()
});
let now = Duration::from_secs(100);
cache.store(&request, json!({"answer": 7}), now)?;
assert_eq!(cache.async_lookup(&request, now).await?, Some(json!({"answer": 7})));
```
For Redis, inject `RedisCache::new(url, ttl, ResponseCacheCodec)` instead. Namespaces are optional and existing namespace prefixes are preserved. Sync operations check out independent connections from a bounded pool, while async callers, including counters and claims, move that blocking work off the executor. The pool skips the checkout PING and instead discards any connection whose command failed
Callers supply Unix time for response freshness. Backend TTL uses its own clock. A read can reject an entry through `max_age` even while the backend still retains it
## Python integration boundary
The extension keeps a private test harness for memory and Redis single and batch response lookup and storage. Batch lookup returns ordered values plus missing indices for embedding partial-hit wiring. No bridge-only cache type is part of the public API
Object responses are written as they are, and every other response shape is written as a serialized string, which is the pair of shapes Python reads. A string on the wire is therefore always a serialized response, so string-valued responses round trip. Typed backends such as memory never pass through the codec
The resolver reads the namespace's `cache` attribute each time it resolves. A captured binding retains the selected service for its operation, including background writes. `None` disables caching. Custom Python cache objects keep their original methods, arguments, returned awaitables, exceptions, and caller-task execution
Python callbacks use the built-in `Cache` API, so a `Cache` subclass works unchanged. A batch lookup takes one original kwargs mapping per request and returns the list of `get_cache` or gathered `async_get_cache` results, while native bindings return `{values, missing_indices}`. A batch store hands the caller's original result to `async_add_cache_pipeline`. `ping` calls `ping`, and a flush goes to the facade's backend
The private facade test harness checks object identity, method overrides, effective TTL, Redis namespace, memory capacity, and later configuration changes before selecting native execution. Its snapshot includes Redis connection settings, so a later `redis_kwargs` change, including an SSL option, selects Python callback execution. Buffered async writes honor `redis_flush_size`. Public activation must construct the shared native service from the initial Python Redis settings, including `litellm.default_redis_ttl` and SSL options. A buffered entry keeps the time it was produced, and a failed flush drops its batch instead of growing the buffer during an outage. The harness does not migrate entries or replace Python methods. Until activation configures one shared service, the Python facade and native test service can hold separate data. Existing public cache constructors remain on Python
Native cache handles must be recreated after fork. The bridge releases the GIL around native operations, and Redis runs blocking connection operations off the async executor. Native errors propagate to the host, which owns the existing fail-open and logging policy
The Redis backend also provides the primitives needed to preserve its direct Python surface later: TLS URLs, ping, bulk delete, counter batches, TTL, scan, set membership, raw queue push and pop, queue and counter pipelines, counter floor and maximum operations, script evaluation, client information, namespaced flush, and full flush. These are backend operations only and are not exported to Python by this PR. Memory provides TTL, oldest-key, and counter-pipeline operations
## Adding another backend
Implement `BaseCache` for the backend with its associated value type, and accept a `CacheCodec` when wire serialization is needed. `ResponseCache<B>` then works without another response implementation. Add a concrete bridge enum variant and constructor only when exposing that backend to Python
Verify typed values, TTL precedence, missing entries, serialization failures, namespaces, batch ordering, and sync/async behavior. Run response fixtures with `ResponseCacheCodec`, including both Python envelope encodings, before enabling a public facade
## Follow-up scope
Public SDK, Router, and proxy activation still need constructor parity, stream replay, embedding partial-batch integration, response reconstruction, callback scheduling, and failure-policy integration. This foundation does not switch those request paths
Redis cluster, disk, cloud stores, and semantic caching remain follow-ups. The generic dual cache takes read, write, and remote-failure policies, runs its async operations through the async L2 methods, and provides L2-first counters and atomic affinity claims. Errors propagate by default, and `RemoteFailurePolicy::UseLocal` opts key-value operations and claims into the local tier when L2 is unavailable. Claims compare decoded values, so a pin written by Python still matches. Public Router integration remains follow-up work. Reservations and pubsub still need explicit capabilities owned by their consuming features. Adding a cache backend does not establish those guarantees

View file

@ -0,0 +1,45 @@
use std::{sync::Mutex, time::Duration};
use litellm_cache::{BaseCache, Error, ExactCacheContext};
use serde_json::Value;
use crate::{CacheEntry, ResponseCache, ResponseCacheRequest};
pub struct WriteBuffer {
flush_size: usize,
entries: Mutex<Vec<(ResponseCacheRequest, Value, Duration)>>,
}
impl WriteBuffer {
pub fn new(flush_size: usize) -> Self {
Self {
flush_size: flush_size.max(1),
entries: Mutex::new(Vec::new()),
}
}
pub async fn async_store<B: BaseCache<Value = CacheEntry, Context = ExactCacheContext>>(
&self,
cache: &ResponseCache<B>,
request: &ResponseCacheRequest,
response: Value,
now: Duration,
) -> Result<(), Error> {
let pending = {
let mut entries = self.entries.lock().map_err(|_| Error::Unavailable)?;
entries.push((request.clone(), response, now));
(entries.len() >= self.flush_size).then(|| std::mem::take(&mut *entries))
};
// A failed flush drops its batch, as Python does. Requeueing would grow the
// buffer and re-send an ever larger pipeline on every write during an outage.
match pending {
Some(pending) => cache.async_store_entries(pending).await,
None => Ok(()),
}
}
pub fn clear(&self) -> Result<(), Error> {
self.entries.lock().map_err(|_| Error::Unavailable)?.clear();
Ok(())
}
}

View file

@ -0,0 +1,147 @@
use std::time::Duration;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sha2::{Digest, Sha256};
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize, PartialEq, Eq)]
pub enum CacheMode {
#[default]
#[serde(rename = "default_on")]
DefaultOn,
#[serde(rename = "default_off")]
DefaultOff,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct CacheKeyField {
pub name: String,
pub value: Option<String>,
pub api_parameter: bool,
pub internal_parameter: bool,
}
#[derive(Clone, Debug, Default, Deserialize, Serialize)]
#[serde(default)]
pub struct CacheKeyInput {
pub fields: Vec<CacheKeyField>,
pub preset: Option<String>,
pub namespace: Option<String>,
pub include_provider_parameters: bool,
}
#[derive(Default)]
pub struct CacheKeyContext {
pub model_group: Option<String>,
pub caching_groups: Vec<(Vec<String>, String)>,
pub file_checksum: Option<String>,
pub file_object_name: Option<String>,
pub metadata_file_name: Option<String>,
pub parameters_file_name: Option<String>,
}
impl CacheKeyContext {
pub fn apply(self, input: &mut CacheKeyInput) {
let group = self.model_group.as_ref().and_then(|model| {
self.caching_groups
.iter()
.find(|(models, _)| models.contains(model))
});
for field in &mut input.fields {
match field.name.as_str() {
"model" => {
field.value = group
.map(|(_, formatted)| formatted.clone())
.or_else(|| self.model_group.clone())
.or_else(|| field.value.take())
}
"file" => {
field.value = self
.file_checksum
.clone()
.or_else(|| self.file_object_name.clone())
.or_else(|| self.metadata_file_name.clone())
.or_else(|| self.parameters_file_name.clone())
}
_ => {}
}
}
}
}
pub fn get_cache_key(input: &CacheKeyInput) -> String {
cache_key(input)
}
pub fn cache_key(input: &CacheKeyInput) -> String {
if let Some(preset) = &input.preset {
return preset.clone();
}
let mut digest = Sha256::new();
for field in &input.fields {
if (field.api_parameter || (input.include_provider_parameters && !field.internal_parameter))
&& let Some(value) = &field.value
{
digest.update(field.name.as_bytes());
digest.update(b": ");
digest.update(value.as_bytes());
}
}
let hash = format!("{:x}", digest.finalize());
input
.namespace
.as_deref()
.filter(|namespace| !namespace.is_empty())
.map_or(hash.clone(), |namespace| format!("{namespace}:{hash}"))
}
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
pub struct CacheControls {
pub supported_call_type: bool,
pub configured: bool,
pub native_backend: bool,
pub default_on: bool,
pub caching: Option<bool>,
pub no_cache: bool,
pub no_store: bool,
#[serde(default)]
pub use_cache: bool,
}
impl CacheControls {
pub fn reads(self) -> bool {
self.supported_call_type
&& self.configured
&& self.caching.unwrap_or(true)
&& !self.no_cache
&& (self.default_on || self.use_cache)
}
pub fn writes(self) -> bool {
self.supported_call_type
&& self.configured
&& self.caching.unwrap_or(true)
&& !self.no_store
&& (self.default_on || self.use_cache)
}
}
pub fn should_use_cache(controls: CacheControls) -> bool {
controls.reads() || controls.writes()
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct CacheEntry {
#[serde(skip_serializing_if = "Option::is_none")]
pub timestamp: Option<f64>,
pub response: Value,
}
impl CacheEntry {
pub fn fresh(&self, now: Duration, max_age: Option<Duration>) -> bool {
self.timestamp.is_none_or(|timestamp| {
timestamp.is_finite()
&& max_age.is_none_or(|age| now.as_secs_f64() - timestamp <= age.as_secs_f64())
})
}
}

View file

@ -0,0 +1,129 @@
use litellm_cache::{CacheCodec, Error};
use serde_json::Value;
use crate::CacheEntry;
#[derive(Clone, Copy, Debug, Default)]
pub struct ResponseCacheCodec;
impl CacheCodec for ResponseCacheCodec {
type Value = CacheEntry;
fn encode(&self, value: &CacheEntry) -> Result<Vec<u8>, Error> {
if value
.timestamp
.is_some_and(|timestamp| !timestamp.is_finite())
{
return Err(Error::InvalidEntry);
}
// Python reads a `response` that is either a dict or a serialized string, so every
// other shape is written serialized. A string on the wire is therefore always a
// serialized response, which keeps string-valued responses unambiguous.
if value.timestamp.is_none() || value.response.is_object() {
return serde_json::to_vec(value).map_err(|_| Error::InvalidEntry);
}
let response = serde_json::to_string(&value.response).map_err(|_| Error::InvalidEntry)?;
serde_json::to_vec(&CacheEntry {
timestamp: value.timestamp,
response: Value::String(response),
})
.map_err(|_| Error::InvalidEntry)
}
fn decode(&self, bytes: &[u8]) -> Result<CacheEntry, Error> {
let text = std::str::from_utf8(bytes).map_err(|_| Error::InvalidEntry)?;
let value = decode_value(text)?;
let Some(timestamp) = value.get("timestamp") else {
return Ok(CacheEntry {
timestamp: None,
response: value,
});
};
let Some(timestamp) = timestamp.as_f64().filter(|timestamp| timestamp.is_finite()) else {
return Err(Error::InvalidEntry);
};
let response = match value.get("response").ok_or(Error::InvalidEntry)? {
Value::String(text) => decode_value(text)?,
response => response.clone(),
};
Ok(CacheEntry {
timestamp: Some(timestamp),
response,
})
}
}
fn decode_value(text: &str) -> Result<Value, Error> {
if let Ok(value) = serde_json::from_str(text) {
return Ok(value);
}
check_literal_depth(text)?;
let literal: py_literal::Value = text.parse().map_err(|_| Error::InvalidEntry)?;
literal_value(literal, 0)
}
fn literal_value(value: py_literal::Value, depth: usize) -> Result<Value, Error> {
use py_literal::Value as Literal;
if depth > 128 {
return Err(Error::InvalidEntry);
}
match value {
Literal::String(text) => Ok(Value::String(text)),
Literal::Boolean(value) => Ok(Value::Bool(value)),
Literal::None => Ok(Value::Null),
Literal::Integer(value) => {
serde_json::from_str(&value.to_string()).map_err(|_| Error::InvalidEntry)
}
Literal::Float(value) => serde_json::Number::from_f64(value)
.map(Value::Number)
.ok_or(Error::InvalidEntry),
Literal::List(values) | Literal::Tuple(values) => values
.into_iter()
.map(|value| literal_value(value, depth + 1))
.collect::<Result<Vec<_>, _>>()
.map(Value::Array),
Literal::Dict(entries) => entries
.into_iter()
.map(|(key, value)| {
let Literal::String(key) = key else {
return Err(Error::InvalidEntry);
};
Ok((key, literal_value(value, depth + 1)?))
})
.collect::<Result<serde_json::Map<_, _>, _>>()
.map(Value::Object),
_ => Err(Error::InvalidEntry),
}
}
fn check_literal_depth(text: &str) -> Result<(), Error> {
let mut quote = None;
let mut escaped = false;
let mut depth = 0usize;
for ch in text.chars() {
if escaped {
escaped = false;
continue;
}
if let Some(delimiter) = quote {
if ch == '\\' {
escaped = true;
} else if ch == delimiter {
quote = None;
}
continue;
}
match ch {
'\'' | '"' => quote = Some(ch),
'[' | '{' | '(' => {
depth += 1;
if depth > 128 {
return Err(Error::InvalidEntry);
}
}
']' | '}' | ')' => depth = depth.saturating_sub(1),
_ => {}
}
}
Ok(())
}

View file

@ -0,0 +1,22 @@
use serde::Serialize;
use serde_json::Value;
#[derive(Clone, Debug, PartialEq, Serialize)]
pub struct PartialHits {
pub values: Vec<Option<Value>>,
pub missing_indices: Vec<usize>,
}
impl PartialHits {
pub fn new(values: Vec<Option<Value>>) -> Self {
let missing_indices = values
.iter()
.enumerate()
.filter_map(|(index, value)| value.is_none().then_some(index))
.collect();
Self {
values,
missing_indices,
}
}
}

View file

@ -0,0 +1,14 @@
mod buffer;
mod caching;
mod codec;
mod embedding;
mod response;
pub use buffer::WriteBuffer;
pub use caching::{
CacheControls, CacheEntry, CacheKeyContext, CacheKeyField, CacheKeyInput, CacheMode, cache_key,
get_cache_key, should_use_cache,
};
pub use codec::ResponseCacheCodec;
pub use embedding::PartialHits;
pub use response::{ResponseCache, ResponseCacheRequest};

View file

@ -0,0 +1,280 @@
use std::{sync::Arc, time::Duration};
use litellm_cache::{
BaseCache, BatchCache, BatchEntry, CacheConnectionResult, Error, ExactCacheContext, FlushCache,
};
use serde_json::Value;
use crate::{CacheControls, CacheEntry, CacheKeyInput, PartialHits, cache_key};
#[derive(Clone)]
pub struct ResponseCacheRequest {
pub key: CacheKeyInput,
pub controls: CacheControls,
pub context: ExactCacheContext,
pub max_age: Option<Duration>,
}
impl ResponseCacheRequest {
pub fn new(key: CacheKeyInput) -> Self {
Self {
key,
controls: CacheControls {
configured: true,
supported_call_type: true,
native_backend: true,
default_on: true,
..Default::default()
},
context: ExactCacheContext::default(),
max_age: None,
}
}
}
pub struct ResponseCache<B: BaseCache<Value = CacheEntry, Context = ExactCacheContext>> {
backend: Arc<B>,
}
impl<B: BaseCache<Value = CacheEntry, Context = ExactCacheContext>> ResponseCache<B> {
pub fn new(backend: Arc<B>) -> Self {
Self { backend }
}
pub fn backend(&self) -> &B {
&self.backend
}
pub fn default_ttl(&self) -> Option<Duration> {
self.backend.get_ttl(&ExactCacheContext::default())
}
pub async fn async_flush(&self) -> Result<(), Error>
where
B: FlushCache,
{
self.backend.async_flush_cache().await
}
pub async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
self.backend.test_connection().await
}
pub fn lookup(
&self,
request: &ResponseCacheRequest,
now: Duration,
) -> Result<Option<Value>, Error> {
if !request.controls.reads() {
return Ok(None);
}
let entry = match self
.backend
.get_cache(&cache_key(&request.key), &request.context)
{
Ok(entry) => entry,
Err(Error::InvalidEntry) => None,
Err(error) => return Err(error),
};
Ok(Self::fresh_or_miss(entry, now, request.max_age))
}
pub async fn async_lookup(
&self,
request: &ResponseCacheRequest,
now: Duration,
) -> Result<Option<Value>, Error> {
if !request.controls.reads() {
return Ok(None);
}
let entry = match self
.backend
.async_get_cache(&cache_key(&request.key), &request.context)
.await
{
Ok(entry) => entry,
Err(Error::InvalidEntry) => None,
Err(error) => return Err(error),
};
Ok(Self::fresh_or_miss(entry, now, request.max_age))
}
pub fn lookup_batch(
&self,
requests: &[ResponseCacheRequest],
now: Duration,
) -> Result<PartialHits, Error>
where
B: BatchCache,
{
let readable = requests
.iter()
.enumerate()
.filter(|(_, request)| request.controls.reads())
.collect::<Vec<_>>();
let keys = readable
.iter()
.map(|(_, request)| cache_key(&request.key))
.collect::<Vec<_>>();
let entries = if let Some((_, request)) = readable.first() {
self.backend.batch_get_cache(&keys, &request.context)?
} else {
Vec::new()
};
Self::partial_hits(requests, readable, entries, now)
}
pub async fn async_lookup_batch(
&self,
requests: &[ResponseCacheRequest],
now: Duration,
) -> Result<PartialHits, Error>
where
B: BatchCache,
{
let readable = requests
.iter()
.enumerate()
.filter(|(_, request)| request.controls.reads())
.collect::<Vec<_>>();
let keys = readable
.iter()
.map(|(_, request)| cache_key(&request.key))
.collect::<Vec<_>>();
let entries = if let Some((_, request)) = readable.first() {
self.backend
.async_batch_get_cache(keys, request.context.clone())
.await?
} else {
Vec::new()
};
Self::partial_hits(requests, readable, entries, now)
}
pub fn store(
&self,
request: &ResponseCacheRequest,
response: Value,
now: Duration,
) -> Result<(), Error> {
if !request.controls.writes() {
return Ok(());
}
self.backend.set_cache(
&cache_key(&request.key),
CacheEntry {
timestamp: Some(now.as_secs_f64()),
response,
},
&request.context,
)
}
pub async fn async_store(
&self,
request: &ResponseCacheRequest,
response: Value,
now: Duration,
) -> Result<(), Error> {
if !request.controls.writes() {
return Ok(());
}
self.backend
.async_set_cache(
&cache_key(&request.key),
CacheEntry {
timestamp: Some(now.as_secs_f64()),
response,
},
request.context.clone(),
)
.await
}
pub async fn async_store_batch(
&self,
entries: Vec<(ResponseCacheRequest, Value)>,
now: Duration,
) -> Result<(), Error> {
self.async_store_entries(
entries
.into_iter()
.map(|(request, response)| (request, response, now))
.collect(),
)
.await
}
/// Stores entries that each carry the time they were produced, so a deferred write keeps
/// the freshness of its original response.
pub async fn async_store_entries(
&self,
entries: Vec<(ResponseCacheRequest, Value, Duration)>,
) -> Result<(), Error> {
let writable = entries
.into_iter()
.filter(|(request, _, _)| request.controls.writes())
.map(|(request, response, now)| {
(
cache_key(&request.key),
CacheEntry {
timestamp: Some(now.as_secs_f64()),
response,
},
request.context,
)
})
.collect::<Vec<_>>();
let Some((_, _, first_kwargs)) = writable.first() else {
return Ok(());
};
if writable
.iter()
.all(|(_, _, context)| context == first_kwargs)
{
let context = first_kwargs.clone();
let cache_list = writable
.into_iter()
.map(|(key, entry, _)| (key, entry))
.collect();
return self
.backend
.async_set_cache_pipeline(cache_list, context)
.await;
}
for (key, entry, context) in writable {
self.backend.async_set_cache(&key, entry, context).await?;
}
Ok(())
}
fn partial_hits(
requests: &[ResponseCacheRequest],
readable: Vec<(usize, &ResponseCacheRequest)>,
entries: Vec<BatchEntry<CacheEntry>>,
now: Duration,
) -> Result<PartialHits, Error> {
if readable.len() != entries.len() {
return Err(Error::Unavailable);
}
let mut values = vec![None; requests.len()];
for ((index, request), entry) in readable.into_iter().zip(entries) {
let response = match entry {
BatchEntry::Hit(entry) => Self::fresh_or_miss(Some(entry), now, request.max_age),
BatchEntry::Miss | BatchEntry::Invalid => None,
};
values[index] = response;
}
Ok(PartialHits::new(values))
}
fn fresh_or_miss(
entry: Option<CacheEntry>,
now: Duration,
max_age: Option<Duration>,
) -> Option<Value> {
entry
.filter(|entry| entry.fresh(now, max_age))
.map(|entry| entry.response)
}
}

View file

@ -0,0 +1,90 @@
use litellm_cache_response::{
CacheControls, CacheKeyContext, CacheKeyField, CacheKeyInput, cache_key, get_cache_key,
};
use sha2::{Digest, Sha256};
#[test]
fn keys_match_python_order_groups_files_presets_and_namespaces() {
let mut input = CacheKeyInput {
fields: vec![
CacheKeyField {
name: "model".into(),
value: Some("deployment".into()),
api_parameter: true,
internal_parameter: false,
},
CacheKeyField {
name: "file".into(),
value: None,
api_parameter: true,
internal_parameter: false,
},
],
namespace: Some("team".into()),
..Default::default()
};
CacheKeyContext {
model_group: Some("group".into()),
caching_groups: vec![(vec!["group".into()], "['group']".into())],
file_checksum: Some("checksum".into()),
..Default::default()
}
.apply(&mut input);
assert_eq!(
cache_key(&input),
format!(
"team:{:x}",
Sha256::digest(b"model: ['group']file: checksum")
)
);
input.preset = Some("preset".into());
assert_eq!(get_cache_key(&input), "preset");
}
#[test]
fn cache_controls_honor_default_modes_and_directives() {
let enabled = CacheControls {
supported_call_type: true,
configured: true,
default_on: true,
..Default::default()
};
assert!(enabled.reads());
assert!(enabled.writes());
assert!(
!CacheControls {
default_on: false,
..enabled
}
.reads()
);
assert!(
CacheControls {
default_on: false,
use_cache: true,
..enabled
}
.reads()
);
assert!(
!CacheControls {
no_cache: true,
..enabled
}
.reads()
);
assert!(
!CacheControls {
no_store: true,
..enabled
}
.writes()
);
assert!(
!CacheControls {
caching: Some(false),
..enabled
}
.writes()
);
}

View file

@ -0,0 +1,484 @@
use std::{
sync::{
Arc,
atomic::{AtomicU64, Ordering},
},
time::Duration,
};
use litellm_cache::{BaseCache, CacheCodec, Error};
use litellm_cache_memory::InMemoryCache;
use litellm_cache_redis::RedisCache;
use litellm_cache_response::{
CacheEntry, CacheKeyField, CacheKeyInput, ResponseCache, ResponseCacheCodec,
ResponseCacheRequest, WriteBuffer,
};
use redis_test::{MockCmd, MockRedisConnection};
use serde_json::json;
fn memory() -> Arc<ResponseCache<InMemoryCache<CacheEntry>>> {
Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
Some(8),
Some(Duration::from_secs(600)),
))))
}
fn request() -> ResponseCacheRequest {
ResponseCacheRequest::new(CacheKeyInput {
preset: Some("tenant:key".into()),
..Default::default()
})
}
#[tokio::test]
async fn sync_and_async_consumers_share_keys_ttls_and_freshness() {
let clock = Arc::new(AtomicU64::new(100));
let backend = Arc::new(InMemoryCache::with_clock(
Some(8),
Some(Duration::from_secs(600)),
{
let clock = clock.clone();
move || Duration::from_secs(clock.load(Ordering::SeqCst))
},
));
let cache = ResponseCache::new(backend.clone());
let mut request = request();
request.context.ttl = Some(Duration::from_secs(10));
request.max_age = Some(Duration::from_secs(5));
cache
.store(
&request,
json!({"choices": [1], "usage": {"total_tokens": 7}}),
Duration::from_secs(100),
)
.unwrap();
assert_eq!(
backend.expires_at("tenant:key").unwrap(),
Some(Duration::from_secs(110))
);
assert!(
cache
.async_lookup(&request, Duration::from_secs(105))
.await
.unwrap()
.is_some()
);
assert_eq!(
cache.lookup(&request, Duration::from_secs(106)).unwrap(),
None
);
request.max_age = None;
assert_eq!(
cache
.lookup(&request, Duration::from_secs(106))
.unwrap()
.unwrap()["usage"]["total_tokens"],
7
);
clock.store(111, Ordering::SeqCst);
assert_eq!(
cache
.async_lookup(&request, Duration::from_secs(111))
.await
.unwrap(),
None
);
cache
.async_store(&request, json!({"choices": [2]}), Duration::from_secs(111))
.await
.unwrap();
assert_eq!(
cache.lookup(&request, Duration::from_secs(111)).unwrap(),
Some(json!({"choices": [2]}))
);
}
#[tokio::test]
async fn directives_skip_io_and_keep_reads_and_writes_independent() {
let cache = memory();
let mut request = request();
let now = Duration::from_secs(100);
request.controls.no_store = true;
cache
.async_store(&request, json!({"v": 1}), now)
.await
.unwrap();
assert_eq!(cache.lookup(&request, now).unwrap(), None);
request.controls.no_store = false;
request.controls.no_cache = true;
cache.store(&request, json!({"v": 2}), now).unwrap();
assert_eq!(cache.async_lookup(&request, now).await.unwrap(), None);
request.controls.no_cache = false;
assert_eq!(cache.lookup(&request, now).unwrap(), Some(json!({"v": 2})));
request.controls.default_on = false;
cache.store(&request, json!({"v": 3}), now).unwrap();
assert_eq!(cache.lookup(&request, now).unwrap(), None);
request.controls.use_cache = true;
assert_eq!(cache.lookup(&request, now).unwrap(), Some(json!({"v": 2})));
request.controls.supported_call_type = false;
assert_eq!(cache.lookup(&request, now).unwrap(), None);
}
#[tokio::test]
async fn redis_consumer_reads_python_sync_and_async_envelopes_and_writes_compatible_json() {
let connection = MockRedisConnection::new([
MockCmd::new(
redis::cmd("GET").arg("tenant:key"),
Ok(br#"{'timestamp': 100.0, 'response': '{"ok": true, "text": "cached"}'}"#.to_vec()),
),
MockCmd::new(
redis::cmd("GET").arg("tenant:key"),
Ok(br#"{"timestamp":100.0,"response":{"ok":true,"text":"cached"}}"#.to_vec()),
),
MockCmd::new(
redis::cmd("SETEX")
.arg("tenant:key")
.arg(600)
.arg(br#"{"timestamp":100.0,"response":{"ok":true,"text":"cached"}}"#.as_slice()),
Ok("OK"),
),
])
.assert_all_commands_consumed();
let backend = RedisCache::with_connection(connection, None, ResponseCacheCodec)
.with_namespace(Some("tenant".into()));
let cache = ResponseCache::new(Arc::new(backend));
let request = request();
let expected = json!({"ok": true, "text": "cached"});
assert_eq!(
cache.lookup(&request, Duration::from_secs(101)).unwrap(),
Some(expected.clone())
);
assert_eq!(
cache
.async_lookup(&request, Duration::from_secs(101))
.await
.unwrap(),
Some(expected.clone())
);
cache
.async_store(&request, expected, Duration::from_secs(100))
.await
.unwrap();
}
#[tokio::test]
async fn captured_service_keeps_the_selected_backend_for_background_writes() {
let original = memory();
let captured = original.clone();
let replacement = memory();
let request = request();
let writer = tokio::spawn({
let request = request.clone();
async move {
captured
.async_store(
&request,
json!({"selected": "original"}),
Duration::from_secs(100),
)
.await
}
});
writer.await.unwrap().unwrap();
assert_eq!(
original.lookup(&request, Duration::from_secs(100)).unwrap(),
Some(json!({"selected":"original"}))
);
assert_eq!(
replacement
.lookup(&request, Duration::from_secs(100))
.unwrap(),
None
);
}
#[test]
fn generated_keys_preserve_namespace_and_explicit_keys() {
let cache = memory();
let key = CacheKeyInput {
fields: vec![CacheKeyField {
name: "model".into(),
value: Some("a".into()),
api_parameter: true,
internal_parameter: false,
}],
namespace: Some("tenant".into()),
..Default::default()
};
let generated = ResponseCacheRequest::new(key.clone());
let explicit = ResponseCacheRequest::new(CacheKeyInput {
preset: Some(litellm_cache_response::cache_key(&key)),
..Default::default()
});
cache
.store(&generated, json!({"value": 7}), Duration::from_secs(100))
.unwrap();
assert_eq!(
cache.lookup(&explicit, Duration::from_secs(100)).unwrap(),
Some(json!({"value":7}))
);
}
#[test]
fn response_codec_accepts_python_literals_without_executing_code() {
let bytes = br#"{'timestamp': 100.0, 'response': {'text': 'hello \\ world', 'flag': True, 'empty': None, 'list': [1, 2.5]}}"#;
let entry = ResponseCacheCodec.decode(bytes).unwrap();
assert_eq!(
entry.response,
json!({"text": "hello \\ world", "flag": true, "empty": null, "list": [1, 2.5]})
);
for bytes in [
b"__import__('os').system('false')".as_slice(),
b"{'timestamp': 'invalid', 'response': {}}",
b"{'timestamp': 1e9999, 'response': {}}",
] {
assert_eq!(
ResponseCacheCodec.decode(bytes).unwrap_err(),
Error::InvalidEntry
);
}
let deep = format!("{}None{}", "[".repeat(1000), "]".repeat(1000));
assert_eq!(
ResponseCacheCodec.decode(deep.as_bytes()).unwrap_err(),
Error::InvalidEntry
);
assert_eq!(
ResponseCacheCodec
.encode(&CacheEntry {
timestamp: Some(f64::NAN),
response: json!({})
})
.unwrap_err(),
Error::InvalidEntry
);
}
#[tokio::test]
async fn invalid_entries_are_misses_and_disabled_reads_do_not_touch_redis() {
let connection = MockRedisConnection::new([MockCmd::new(
redis::cmd("GET").arg("tenant:key"),
Ok(b"invalid".to_vec()),
)])
.assert_all_commands_consumed();
let backend = RedisCache::with_connection(connection, None, ResponseCacheCodec);
let cache = ResponseCache::new(Arc::new(backend));
let mut request = request();
request.controls.no_cache = true;
assert_eq!(cache.lookup(&request, Duration::ZERO).unwrap(), None);
request.controls.no_cache = false;
assert_eq!(
cache.async_lookup(&request, Duration::ZERO).await.unwrap(),
None
);
}
#[test]
fn string_responses_round_trip_through_typed_and_wire_backends() {
let cache = ResponseCache::new(Arc::new(InMemoryCache::default()));
let now = Duration::from_secs(100);
for response in [json!("hello world"), json!("123"), json!("null")] {
cache.store(&request(), response.clone(), now).unwrap();
assert_eq!(
cache.lookup(&request(), now).unwrap(),
Some(response.clone())
);
let wire = ResponseCacheCodec
.encode(&CacheEntry {
timestamp: Some(100.0),
response: response.clone(),
})
.unwrap();
assert_eq!(ResponseCacheCodec.decode(&wire).unwrap().response, response);
}
}
#[test]
fn non_object_responses_are_written_as_python_readable_serialized_strings() {
let wire = ResponseCacheCodec
.encode(&CacheEntry {
timestamp: Some(100.0),
response: json!([1, 2]),
})
.unwrap();
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&wire).unwrap(),
json!({"timestamp": 100.0, "response": "[1,2]"})
);
assert_eq!(
ResponseCacheCodec.decode(&wire).unwrap().response,
json!([1, 2])
);
assert_eq!(
ResponseCacheCodec.decode(br#"{"timestamp": 100.0, "response": "not serialized"}"#),
Err(Error::InvalidEntry)
);
}
#[test]
fn response_entries_preserve_the_existing_json_representation() {
let codec = ResponseCacheCodec;
let entry = CacheEntry {
timestamp: Some(123.0),
response: json!({"choices": [{"text": "cached"}]}),
};
let bytes = codec.encode(&entry).unwrap();
assert_eq!(bytes, serde_json::to_vec(&entry).unwrap());
assert_eq!(codec.decode(&bytes).unwrap(), entry);
}
#[test]
fn response_codec_preserves_values_without_timestamps() {
let codec = ResponseCacheCodec;
let raw = json!({"choices": [{"text": "legacy"}]});
let entry = codec.decode(&serde_json::to_vec(&raw).unwrap()).unwrap();
assert_eq!(entry.timestamp, None);
assert_eq!(entry.response, raw);
let backend = Arc::new(InMemoryCache::default());
BaseCache::set_cache(backend.as_ref(), "tenant:key", entry, &Default::default()).unwrap();
let cache = ResponseCache::new(backend);
assert_eq!(
cache.lookup(&request(), Duration::from_secs(100)).unwrap(),
Some(json!({"choices": [{"text": "legacy"}]}))
);
}
#[tokio::test]
async fn batch_lookup_reports_partial_hits_and_batch_store_populates_misses() {
let cache = memory();
let requests = ["hit", "miss", "disabled"].map(|key| {
ResponseCacheRequest::new(CacheKeyInput {
preset: Some(key.into()),
..Default::default()
})
});
cache
.store(&requests[0], json!({"value": 1}), Duration::from_secs(100))
.unwrap();
let mut requests = requests.to_vec();
requests[2].controls.caching = Some(false);
let partial = cache
.async_lookup_batch(&requests, Duration::from_secs(100))
.await
.unwrap();
assert_eq!(partial.values, vec![Some(json!({"value": 1})), None, None]);
assert_eq!(partial.missing_indices, vec![1, 2]);
cache
.async_store_batch(
vec![
(requests[1].clone(), json!({"value": 2})),
(requests[2].clone(), json!({"value": 3})),
],
Duration::from_secs(100),
)
.await
.unwrap();
assert_eq!(
cache
.lookup(&requests[1], Duration::from_secs(100))
.unwrap(),
Some(json!({"value": 2}))
);
requests[2].controls.caching = None;
assert_eq!(
cache
.lookup(&requests[2], Duration::from_secs(100))
.unwrap(),
None
);
}
#[tokio::test]
async fn deferred_entries_keep_the_time_they_were_produced() {
let cache = ResponseCache::new(Arc::new(InMemoryCache::default()));
let mut request = request();
request.max_age = Some(Duration::from_secs(10));
cache
.async_store_entries(vec![(
request.clone(),
json!({"answer": 7}),
Duration::from_secs(100),
)])
.await
.unwrap();
assert_eq!(
cache.lookup(&request, Duration::from_secs(110)).unwrap(),
Some(json!({"answer": 7}))
);
assert_eq!(
cache.lookup(&request, Duration::from_secs(111)).unwrap(),
None
);
}
#[tokio::test]
async fn write_buffer_flushes_at_its_size_and_keeps_each_produced_time() {
let cache = ResponseCache::new(Arc::new(InMemoryCache::default()));
let buffer = WriteBuffer::new(2);
let mut first = request();
first.max_age = Some(Duration::from_secs(10));
let mut second = request();
second.key.preset = Some("tenant:other".into());
buffer
.async_store(
&cache,
&first,
json!({"answer": 7}),
Duration::from_secs(100),
)
.await
.unwrap();
assert_eq!(
cache.lookup(&first, Duration::from_secs(100)).unwrap(),
None
);
buffer
.async_store(
&cache,
&second,
json!({"answer": 8}),
Duration::from_secs(200),
)
.await
.unwrap();
assert_eq!(
cache.lookup(&first, Duration::from_secs(110)).unwrap(),
Some(json!({"answer": 7}))
);
assert_eq!(
cache.lookup(&first, Duration::from_secs(111)).unwrap(),
None
);
assert_eq!(
cache.lookup(&second, Duration::from_secs(200)).unwrap(),
Some(json!({"answer": 8}))
);
}
#[tokio::test]
async fn write_buffer_clear_drops_pending_entries() {
let cache = ResponseCache::new(Arc::new(InMemoryCache::default()));
let buffer = WriteBuffer::new(2);
let mut other = request();
other.key.preset = Some("tenant:other".into());
let now = Duration::from_secs(100);
buffer
.async_store(&cache, &request(), json!({"answer": 7}), now)
.await
.unwrap();
buffer.clear().unwrap();
buffer
.async_store(&cache, &other, json!({"answer": 8}), now)
.await
.unwrap();
assert_eq!(cache.lookup(&request(), now).unwrap(), None);
assert_eq!(cache.lookup(&other, now).unwrap(), None);
}

View file

@ -8,8 +8,8 @@ repository.workspace = true
[dependencies]
serde.workspace = true
serde_json.workspace = true
sha2.workspace = true
thiserror.workspace = true
[dev-dependencies]
rstest.workspace = true
tokio.workspace = true

View file

@ -1,18 +1,35 @@
use std::future::Future;
use std::pin::Pin;
use std::time::Duration;
use std::{future::Future, time::Duration};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use crate::Error;
pub type CacheFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
#[derive(Clone, Debug, PartialEq)]
pub enum BatchEntry<V> {
Hit(V),
Miss,
Invalid,
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct CacheKwargs {
pub trait CacheContext: Clone + Send + Sync + 'static {
fn ttl(&self) -> Option<Duration>;
fn with_ttl(&self, ttl: Option<Duration>) -> Self;
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct ExactCacheContext {
pub ttl: Option<Duration>,
pub extras: Map<String, Value>,
}
impl CacheContext for ExactCacheContext {
fn ttl(&self) -> Option<Duration> {
self.ttl
}
fn with_ttl(&self, ttl: Option<Duration>) -> Self {
Self { ttl }
}
}
#[derive(Clone, Copy, Debug, Deserialize, Serialize, PartialEq, Eq)]
@ -32,67 +49,59 @@ pub struct CacheConnectionResult {
pub trait BaseCache: Send + Sync {
type Value: Clone + Send + Sync + 'static;
type Context: CacheContext;
fn default_ttl(&self) -> Duration {
Duration::from_secs(60)
}
fn get_ttl(&self, context: &Self::Context) -> Option<Duration>;
fn get_ttl(&self, kwargs: &CacheKwargs) -> Duration {
kwargs.ttl.unwrap_or_else(|| self.default_ttl())
}
fn set_cache(&self, key: &str, value: Self::Value, kwargs: CacheKwargs) -> Result<(), Error>;
fn get_cache(&self, key: &str, kwargs: &CacheKwargs) -> Result<Option<Self::Value>, Error>;
fn async_set_cache<'a>(
&'a self,
key: &'a str,
fn set_cache(
&self,
key: &str,
value: Self::Value,
kwargs: CacheKwargs,
) -> CacheFuture<'a, ()> {
Box::pin(async move { self.set_cache(key, value, kwargs) })
context: &Self::Context,
) -> Result<(), Error>;
fn get_cache(&self, key: &str, context: &Self::Context) -> Result<Option<Self::Value>, Error>;
fn async_set_cache(
&self,
key: &str,
value: Self::Value,
context: Self::Context,
) -> impl Future<Output = Result<(), Error>> + Send {
async move { self.set_cache(key, value, &context) }
}
fn async_get_cache<'a>(
&'a self,
key: &'a str,
kwargs: &'a CacheKwargs,
) -> CacheFuture<'a, Option<Self::Value>> {
Box::pin(async move { self.get_cache(key, kwargs) })
fn async_get_cache(
&self,
key: &str,
context: &Self::Context,
) -> impl Future<Output = Result<Option<Self::Value>, Error>> + Send {
async move { self.get_cache(key, context) }
}
fn async_set_cache_pipeline<'a>(
&'a self,
cache_list: Vec<(String, Self::Value)>,
kwargs: CacheKwargs,
) -> CacheFuture<'a, ()> {
Box::pin(async move {
for (key, value) in cache_list {
self.set_cache(&key, value, kwargs.clone())?;
fn async_set_cache_pipeline(
&self,
entries: Vec<(String, Self::Value)>,
context: Self::Context,
) -> impl Future<Output = Result<(), Error>> + Send {
async move {
for (key, value) in entries {
self.async_set_cache(&key, value, context.clone()).await?;
}
Ok(())
})
}
}
fn batch_cache_write<'a>(
&'a self,
key: &'a str,
fn batch_cache_write(
&self,
key: &str,
value: Self::Value,
kwargs: CacheKwargs,
) -> CacheFuture<'a, ()> {
self.async_set_cache(key, value, kwargs)
context: Self::Context,
) -> impl Future<Output = Result<(), Error>> + Send {
self.async_set_cache(key, value, context)
}
fn delete_cache(&self, key: &str) -> Result<(), Error>;
fn disconnect(&self) -> impl Future<Output = Result<(), Error>> + Send;
fn async_delete_cache<'a>(&'a self, key: &'a str) -> CacheFuture<'a, ()> {
Box::pin(async move { self.delete_cache(key) })
}
fn flush_cache(&self) -> Result<(), Error>;
fn disconnect(&self) -> CacheFuture<'_, ()>;
fn test_connection(&self) -> CacheFuture<'_, CacheConnectionResult>;
fn test_connection(&self) -> impl Future<Output = Result<CacheConnectionResult, Error>> + Send;
}

View file

@ -0,0 +1,85 @@
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Debug, Deserialize, Serialize, PartialEq, Eq, Hash)]
pub enum CacheType {
#[serde(rename = "local")]
Local,
#[serde(rename = "redis")]
Redis,
#[serde(rename = "redis-semantic")]
RedisSemantic,
#[serde(rename = "valkey-semantic")]
ValkeySemantic,
#[serde(rename = "s3")]
S3,
#[serde(rename = "disk")]
Disk,
#[serde(rename = "qdrant-semantic")]
QdrantSemantic,
#[serde(rename = "azure-blob")]
AzureBlob,
#[serde(rename = "gcs")]
Gcs,
}
impl CacheType {
pub const ALL: [Self; 9] = [
Self::Local,
Self::Redis,
Self::RedisSemantic,
Self::ValkeySemantic,
Self::S3,
Self::Disk,
Self::QdrantSemantic,
Self::AzureBlob,
Self::Gcs,
];
pub const fn as_python_name(self) -> &'static str {
match self {
Self::Local => "local",
Self::Redis => "redis",
Self::RedisSemantic => "redis-semantic",
Self::ValkeySemantic => "valkey-semantic",
Self::S3 => "s3",
Self::Disk => "disk",
Self::QdrantSemantic => "qdrant-semantic",
Self::AzureBlob => "azure-blob",
Self::Gcs => "gcs",
}
}
pub fn from_python_name(value: &str) -> Option<Self> {
Self::ALL
.into_iter()
.find(|cache_type| cache_type.as_python_name() == value)
}
}
#[cfg(test)]
mod tests {
use super::CacheType;
#[test]
fn every_python_cache_type_has_one_round_trip_identity() {
let names = CacheType::ALL.map(CacheType::as_python_name);
assert_eq!(
names,
[
"local",
"redis",
"redis-semantic",
"valkey-semantic",
"s3",
"disk",
"qdrant-semantic",
"azure-blob",
"gcs",
]
);
assert_eq!(
names.map(CacheType::from_python_name),
CacheType::ALL.map(Some)
);
}
}

View file

@ -1,166 +1,23 @@
use std::sync::Arc;
use std::time::Duration;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sha2::{Digest, Sha256};
use crate::{BaseCache, CacheKwargs, Error};
pub use crate::BaseCache as Cache;
use crate::{BaseCache, Error};
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize, PartialEq, Eq)]
pub enum CacheMode {
#[default]
#[serde(rename = "default_on")]
DefaultOn,
#[serde(rename = "default_off")]
DefaultOff,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct CacheKeyField {
pub name: String,
pub value: Option<String>,
pub api_parameter: bool,
pub internal_parameter: bool,
}
#[derive(Clone, Debug, Default, Deserialize, Serialize)]
pub struct CacheKeyInput {
pub fields: Vec<CacheKeyField>,
pub preset: Option<String>,
pub namespace: Option<String>,
pub include_provider_parameters: bool,
}
#[derive(Default)]
pub struct CacheKeyContext {
pub model_group: Option<String>,
pub caching_groups: Vec<(Vec<String>, String)>,
pub file_checksum: Option<String>,
pub file_object_name: Option<String>,
pub metadata_file_name: Option<String>,
pub parameters_file_name: Option<String>,
}
impl CacheKeyContext {
pub fn apply(self, input: &mut CacheKeyInput) {
let group = self.model_group.as_ref().and_then(|model| {
self.caching_groups
.iter()
.find(|(models, _)| models.contains(model))
});
for field in &mut input.fields {
match field.name.as_str() {
"model" => {
field.value = group
.map(|(_, formatted)| formatted.clone())
.or_else(|| self.model_group.clone())
.or_else(|| field.value.take())
}
"file" => {
field.value = self
.file_checksum
.clone()
.or_else(|| self.file_object_name.clone())
.or_else(|| self.metadata_file_name.clone())
.or_else(|| self.parameters_file_name.clone())
}
_ => {}
}
}
}
}
pub fn get_cache_key(input: &CacheKeyInput) -> String {
cache_key(input)
}
pub fn cache_key(input: &CacheKeyInput) -> String {
if let Some(preset) = &input.preset {
return preset.clone();
}
let mut digest = Sha256::new();
for field in &input.fields {
if (field.api_parameter || (input.include_provider_parameters && !field.internal_parameter))
&& let Some(value) = &field.value
{
digest.update(field.name.as_bytes());
digest.update(b": ");
digest.update(value.as_bytes());
}
}
let hash = format!("{:x}", digest.finalize());
input
.namespace
.as_deref()
.filter(|namespace| !namespace.is_empty())
.map_or(hash.clone(), |namespace| format!("{namespace}:{hash}"))
}
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
pub struct CacheControls {
pub supported_call_type: bool,
pub configured: bool,
pub native_backend: bool,
pub default_on: bool,
pub caching: Option<bool>,
pub no_cache: bool,
pub no_store: bool,
#[serde(default)]
pub use_cache: bool,
}
impl CacheControls {
pub fn reads(self) -> bool {
self.supported_call_type
&& self.configured
&& self.caching.unwrap_or(true)
&& !self.no_cache
&& (self.default_on || self.use_cache)
}
pub fn writes(self) -> bool {
self.supported_call_type
&& self.configured
&& !self.no_store
&& (self.default_on || self.use_cache)
}
}
pub fn should_use_cache(controls: CacheControls) -> bool {
controls.reads() || controls.writes()
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct CacheEntry {
pub timestamp: f64,
pub response: Value,
}
impl CacheEntry {
pub fn fresh(&self, now: Duration, max_age: Option<Duration>) -> bool {
self.timestamp.is_finite()
&& max_age.is_none_or(|age| now.as_secs_f64() - self.timestamp <= age.as_secs_f64())
}
}
pub fn get_cache(
cache: &dyn BaseCache<Value = CacheEntry>,
pub fn get_cache<B: BaseCache>(
cache: &B,
key: &str,
kwargs: &CacheKwargs,
) -> Result<Option<CacheEntry>, Error> {
cache.get_cache(key, kwargs)
context: &B::Context,
) -> Result<Option<B::Value>, Error> {
cache.get_cache(key, context)
}
pub fn set_cache(
cache: &dyn BaseCache<Value = CacheEntry>,
pub fn set_cache<B: BaseCache>(
cache: &B,
key: &str,
entry: CacheEntry,
kwargs: CacheKwargs,
value: B::Value,
context: &B::Context,
) -> Result<(), Error> {
cache.set_cache(key, entry, kwargs)
cache.set_cache(key, value, context)
}
pub type CacheBackend = Arc<dyn BaseCache<Value = CacheEntry>>;
pub type CacheBackend<B> = Arc<B>;

View file

@ -0,0 +1,169 @@
use std::{future::Future, time::Duration};
use crate::{BaseCache, BatchEntry, Error};
#[derive(Clone, Debug, PartialEq)]
pub struct IncrementOperation {
pub key: String,
pub amount: f64,
pub ttl: Option<Duration>,
}
pub trait BatchCache: BaseCache {
fn batch_get_cache(
&self,
keys: &[String],
context: &Self::Context,
) -> Result<Vec<BatchEntry<Self::Value>>, Error> {
keys.iter()
.map(|key| match self.get_cache(key, context) {
Ok(Some(value)) => Ok(BatchEntry::Hit(value)),
Ok(None) => Ok(BatchEntry::Miss),
Err(Error::InvalidEntry) => Ok(BatchEntry::Invalid),
Err(error) => Err(error),
})
.collect()
}
fn async_batch_get_cache(
&self,
keys: Vec<String>,
context: Self::Context,
) -> impl Future<Output = Result<Vec<BatchEntry<Self::Value>>, Error>> + Send {
async move {
let mut entries = Vec::with_capacity(keys.len());
for key in keys {
entries.push(match self.async_get_cache(&key, &context).await {
Ok(Some(value)) => BatchEntry::Hit(value),
Ok(None) => BatchEntry::Miss,
Err(Error::InvalidEntry) => BatchEntry::Invalid,
Err(error) => return Err(error),
});
}
Ok(entries)
}
}
}
pub trait DeleteCache: BaseCache {
fn delete_cache(&self, key: &str) -> Result<(), Error>;
fn async_delete_cache(&self, key: &str) -> impl Future<Output = Result<(), Error>> + Send {
async move { self.delete_cache(key) }
}
}
pub trait FlushCache: BaseCache {
fn flush_cache(&self) -> Result<(), Error>;
fn async_flush_cache(&self) -> impl Future<Output = Result<(), Error>> + Send {
async move { self.flush_cache() }
}
}
pub trait CounterCache: BaseCache<Value = f64> {
fn increment_cache(&self, key: &str, amount: f64, context: Self::Context)
-> Result<f64, Error>;
fn async_increment(
&self,
key: &str,
amount: f64,
context: Self::Context,
) -> impl Future<Output = Result<f64, Error>> + Send {
async move { self.increment_cache(key, amount, context) }
}
}
pub trait ClaimCache: BaseCache
where
Self::Value: PartialEq,
{
fn claim_cache(
&self,
key: &str,
candidate: Self::Value,
eligible: &[Self::Value],
context: Self::Context,
) -> Result<Self::Value, Error>;
fn async_claim_cache(
&self,
key: &str,
candidate: Self::Value,
eligible: Vec<Self::Value>,
context: Self::Context,
) -> impl Future<Output = Result<Self::Value, Error>> + Send {
async move { self.claim_cache(key, candidate, &eligible, context) }
}
}
pub trait TtlCache: BaseCache {
fn async_get_ttl(
&self,
key: &str,
) -> impl Future<Output = Result<Option<Duration>, Error>> + Send;
}
pub trait SetCache: BaseCache {
type SetValue: Clone + Send + Sync + 'static;
type SetResult: Send + Sync + 'static;
fn async_set_cache_sadd(
&self,
key: &str,
values: Vec<Self::SetValue>,
ttl: Option<Duration>,
) -> impl Future<Output = Result<Self::SetResult, Error>> + Send;
}
pub trait QueueCache: BaseCache {
type QueueValue: Clone + Send + Sync + 'static;
type PopResult: Send + Sync + 'static;
fn async_rpush(
&self,
key: &str,
values: Vec<Self::QueueValue>,
) -> impl Future<Output = Result<usize, Error>> + Send;
fn async_lpop(
&self,
key: &str,
count: Option<usize>,
) -> impl Future<Output = Result<Self::PopResult, Error>> + Send;
}
pub trait ScanCache: BaseCache {
fn async_scan_iter(
&self,
pattern: &str,
count: usize,
) -> impl Future<Output = Result<Vec<String>, Error>> + Send;
}
pub trait ClientInfoCache: BaseCache {
type ClientList: Send + Sync + 'static;
type Info: Send + Sync + 'static;
fn client_list(&self) -> Result<Self::ClientList, Error>;
fn info(&self) -> Result<Self::Info, Error>;
}
pub trait CacheScript: Send + Sync + 'static {
type Argument: Clone + Send + Sync + 'static;
type Output: Send + Sync + 'static;
fn invoke(
&self,
keys: Vec<String>,
arguments: Vec<Self::Argument>,
) -> impl Future<Output = Result<Self::Output, Error>> + Send;
}
pub trait ScriptCache: BaseCache {
type Script: CacheScript;
fn async_register_script(&self, source: String) -> Self::Script;
}

50
litellm-rust/crates/cache/src/codec.rs vendored Normal file
View file

@ -0,0 +1,50 @@
use std::marker::PhantomData;
use serde::{Serialize, de::DeserializeOwned};
use crate::Error;
pub trait CacheCodec: Send + Sync {
type Value: Clone + Send + Sync + 'static;
fn encode(&self, value: &Self::Value) -> Result<Vec<u8>, Error>;
fn decode(&self, bytes: &[u8]) -> Result<Self::Value, Error>;
}
pub struct JsonCodec<V>(PhantomData<fn() -> V>);
impl<V> Clone for JsonCodec<V> {
fn clone(&self) -> Self {
*self
}
}
impl<V> Copy for JsonCodec<V> {}
impl<V> Default for JsonCodec<V> {
fn default() -> Self {
Self::new()
}
}
impl<V> JsonCodec<V> {
pub const fn new() -> Self {
Self(PhantomData)
}
}
impl<V> CacheCodec for JsonCodec<V>
where
V: Clone + Send + Sync + Serialize + DeserializeOwned + 'static,
{
type Value = V;
fn encode(&self, value: &Self::Value) -> Result<Vec<u8>, Error> {
serde_json::to_vec(value).map_err(|_| Error::InvalidEntry)
}
fn decode(&self, bytes: &[u8]) -> Result<Self::Value, Error> {
serde_json::from_slice(bytes).map_err(|_| Error::InvalidEntry)
}
}

390
litellm-rust/crates/cache/src/dual.rs vendored Normal file
View file

@ -0,0 +1,390 @@
use std::{sync::Arc, time::Duration};
use crate::{
BaseCache, BatchCache, BatchEntry, CacheConnectionResult, CacheContext, ClaimCache,
CounterCache, DeleteCache, Error, FlushCache,
};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum ReadPolicy {
#[default]
LocalThenRemote,
LocalOnly,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum WritePolicy {
#[default]
Both,
LocalOnly,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum RemoteFailurePolicy {
#[default]
Propagate,
UseLocal,
}
pub struct DualCache<L1, L2> {
l1: Arc<L1>,
l2: Arc<L2>,
read_policy: ReadPolicy,
write_policy: WritePolicy,
remote_failure_policy: RemoteFailurePolicy,
promotion_ttl: Option<Duration>,
}
impl<L1, L2> DualCache<L1, L2> {
pub fn new(l1: Arc<L1>, l2: Arc<L2>) -> Self {
Self {
l1,
l2,
read_policy: ReadPolicy::default(),
write_policy: WritePolicy::default(),
remote_failure_policy: RemoteFailurePolicy::default(),
promotion_ttl: None,
}
}
pub fn with_read_policy(self, read_policy: ReadPolicy) -> Self {
Self {
read_policy,
..self
}
}
pub fn with_write_policy(self, write_policy: WritePolicy) -> Self {
Self {
write_policy,
..self
}
}
pub fn with_remote_failure_policy(self, remote_failure_policy: RemoteFailurePolicy) -> Self {
Self {
remote_failure_policy,
..self
}
}
pub fn with_promotion_ttl(self, promotion_ttl: Duration) -> Self {
Self {
promotion_ttl: Some(promotion_ttl),
..self
}
}
fn reads_remote(&self) -> bool {
self.read_policy == ReadPolicy::LocalThenRemote
}
fn writes_remote(&self) -> bool {
self.write_policy == WritePolicy::Both
}
fn remote<T>(&self, result: Result<T, Error>) -> Result<Option<T>, Error> {
match result {
Ok(value) => Ok(Some(value)),
Err(Error::Unavailable)
if self.remote_failure_policy == RemoteFailurePolicy::UseLocal =>
{
Ok(None)
}
Err(error) => Err(error),
}
}
fn promotion_context<C: CacheContext>(&self, context: &C) -> C {
context.with_ttl(self.promotion_ttl.or(context.ttl()))
}
}
impl<V, C, L1, L2> DualCache<L1, L2>
where
V: Clone + Send + Sync + 'static,
C: CacheContext,
L1: BaseCache<Value = V, Context = C>,
L2: BaseCache<Value = V, Context = C>,
{
fn missing(entries: &[BatchEntry<V>]) -> Vec<usize> {
entries
.iter()
.enumerate()
.filter_map(|(index, entry)| (!matches!(entry, BatchEntry::Hit(_))).then_some(index))
.collect()
}
fn merge_batch(
&self,
keys: &[String],
context: &C,
mut entries: Vec<BatchEntry<V>>,
missing: Vec<usize>,
remote: Vec<BatchEntry<V>>,
) -> Result<Vec<BatchEntry<V>>, Error> {
if missing.len() != remote.len() {
return Err(Error::Unavailable);
}
for (index, entry) in missing.into_iter().zip(remote) {
if let BatchEntry::Hit(value) = &entry {
let promotion_context = self.promotion_context(context);
self.l1
.set_cache(&keys[index], value.clone(), &promotion_context)?;
}
entries[index] = entry;
}
Ok(entries)
}
}
impl<V, C, L1, L2> BaseCache for DualCache<L1, L2>
where
V: Clone + Send + Sync + 'static,
C: CacheContext,
L1: BaseCache<Value = V, Context = C>,
L2: BaseCache<Value = V, Context = C>,
{
type Value = V;
type Context = C;
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
self.l2.get_ttl(context)
}
fn set_cache(&self, key: &str, value: V, context: &C) -> Result<(), Error> {
if self.writes_remote() {
self.remote(self.l2.set_cache(key, value.clone(), context))?;
}
self.l1.set_cache(key, value, context)
}
fn get_cache(&self, key: &str, context: &C) -> Result<Option<V>, Error> {
if let Some(value) = self.l1.get_cache(key, context)? {
return Ok(Some(value));
}
if !self.reads_remote() {
return Ok(None);
}
let value = self.remote(self.l2.get_cache(key, context))?.flatten();
if let Some(value) = &value {
let promotion_context = self.promotion_context(context);
self.l1.set_cache(key, value.clone(), &promotion_context)?;
}
Ok(value)
}
async fn async_set_cache(&self, key: &str, value: V, context: C) -> Result<(), Error> {
if self.writes_remote() {
self.remote(
self.l2
.async_set_cache(key, value.clone(), context.clone())
.await,
)?;
}
self.l1.async_set_cache(key, value, context).await
}
async fn async_get_cache(&self, key: &str, context: &C) -> Result<Option<V>, Error> {
if let Some(value) = self.l1.async_get_cache(key, context).await? {
return Ok(Some(value));
}
if !self.reads_remote() {
return Ok(None);
}
let value = self
.remote(self.l2.async_get_cache(key, context).await)?
.flatten();
if let Some(value) = &value {
self.l1
.async_set_cache(key, value.clone(), self.promotion_context(context))
.await?;
}
Ok(value)
}
async fn async_set_cache_pipeline(
&self,
entries: Vec<(String, V)>,
context: C,
) -> Result<(), Error> {
if self.writes_remote() {
self.remote(
self.l2
.async_set_cache_pipeline(entries.clone(), context.clone())
.await,
)?;
}
self.l1.async_set_cache_pipeline(entries, context).await
}
async fn disconnect(&self) -> Result<(), Error> {
self.l2.disconnect().await?;
self.l1.disconnect().await
}
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
self.l2.test_connection().await
}
}
impl<V, C, L1, L2> BatchCache for DualCache<L1, L2>
where
V: Clone + Send + Sync + 'static,
C: CacheContext,
L1: BatchCache<Value = V, Context = C>,
L2: BatchCache<Value = V, Context = C>,
{
fn batch_get_cache(&self, keys: &[String], context: &C) -> Result<Vec<BatchEntry<V>>, Error> {
let entries = self.l1.batch_get_cache(keys, context)?;
let missing = Self::missing(&entries);
if missing.is_empty() || !self.reads_remote() {
return Ok(entries);
}
let remote_keys = missing
.iter()
.map(|index| keys[*index].clone())
.collect::<Vec<_>>();
match self.remote(self.l2.batch_get_cache(&remote_keys, context))? {
Some(remote) => self.merge_batch(keys, context, entries, missing, remote),
None => Ok(entries),
}
}
async fn async_batch_get_cache(
&self,
keys: Vec<String>,
context: C,
) -> Result<Vec<BatchEntry<V>>, Error> {
let entries = self
.l1
.async_batch_get_cache(keys.clone(), context.clone())
.await?;
let missing = Self::missing(&entries);
if missing.is_empty() || !self.reads_remote() {
return Ok(entries);
}
let remote_keys = missing.iter().map(|index| keys[*index].clone()).collect();
match self.remote(
self.l2
.async_batch_get_cache(remote_keys, context.clone())
.await,
)? {
Some(remote) => self.merge_batch(&keys, &context, entries, missing, remote),
None => Ok(entries),
}
}
}
impl<V, C, L1, L2> DeleteCache for DualCache<L1, L2>
where
V: Clone + Send + Sync + 'static,
C: CacheContext,
L1: DeleteCache<Value = V, Context = C>,
L2: DeleteCache<Value = V, Context = C>,
{
fn delete_cache(&self, key: &str) -> Result<(), Error> {
if self.writes_remote() {
self.remote(self.l2.delete_cache(key))?;
}
self.l1.delete_cache(key)
}
async fn async_delete_cache(&self, key: &str) -> Result<(), Error> {
if self.writes_remote() {
self.remote(self.l2.async_delete_cache(key).await)?;
}
self.l1.async_delete_cache(key).await
}
}
impl<V, C, L1, L2> FlushCache for DualCache<L1, L2>
where
V: Clone + Send + Sync + 'static,
C: CacheContext,
L1: FlushCache<Value = V, Context = C>,
L2: FlushCache<Value = V, Context = C>,
{
fn flush_cache(&self) -> Result<(), Error> {
if self.writes_remote() {
self.remote(self.l2.flush_cache())?;
}
self.l1.flush_cache()
}
async fn async_flush_cache(&self) -> Result<(), Error> {
if self.writes_remote() {
self.remote(self.l2.async_flush_cache().await)?;
}
self.l1.async_flush_cache().await
}
}
impl<C, L1, L2> CounterCache for DualCache<L1, L2>
where
C: CacheContext,
L1: BaseCache<Value = f64, Context = C>,
L2: CounterCache<Context = C>,
{
fn increment_cache(&self, key: &str, amount: f64, context: C) -> Result<f64, Error> {
let value = self.l2.increment_cache(key, amount, context.clone())?;
self.l1.set_cache(key, value, &context)?;
Ok(value)
}
async fn async_increment(&self, key: &str, amount: f64, context: C) -> Result<f64, Error> {
let value = self
.l2
.async_increment(key, amount, context.clone())
.await?;
self.l1.async_set_cache(key, value, context).await?;
Ok(value)
}
}
impl<V, C, L1, L2> ClaimCache for DualCache<L1, L2>
where
V: Clone + PartialEq + Send + Sync + 'static,
C: CacheContext,
L1: ClaimCache<Value = V, Context = C>,
L2: ClaimCache<Value = V, Context = C>,
{
fn claim_cache(&self, key: &str, candidate: V, eligible: &[V], context: C) -> Result<V, Error> {
match self.remote(
self.l2
.claim_cache(key, candidate.clone(), eligible, context.clone()),
)? {
Some(winner) => {
self.l1.set_cache(key, winner.clone(), &context)?;
Ok(winner)
}
None => self.l1.claim_cache(key, candidate, eligible, context),
}
}
async fn async_claim_cache(
&self,
key: &str,
candidate: V,
eligible: Vec<V>,
context: C,
) -> Result<V, Error> {
match self.remote(
self.l2
.async_claim_cache(key, candidate.clone(), eligible.clone(), context.clone())
.await,
)? {
Some(winner) => {
self.l1
.async_set_cache(key, winner.clone(), context)
.await?;
Ok(winner)
}
None => {
self.l1
.async_claim_cache(key, candidate, eligible, context)
.await
}
}
}
}

View file

@ -4,4 +4,6 @@ pub enum Error {
Unavailable,
#[error("invalid cache entry")]
InvalidEntry,
#[error("flushing Redis requires an explicit namespace")]
UnscopedFlush,
}

View file

@ -1,12 +1,21 @@
mod base_cache;
mod cache_type;
mod caching;
mod capabilities;
mod codec;
mod dual;
mod error;
pub use base_cache::{
BaseCache, CacheConnectionResult, CacheConnectionStatus, CacheFuture, CacheKwargs,
BaseCache, BatchEntry, CacheConnectionResult, CacheConnectionStatus, CacheContext,
ExactCacheContext,
};
pub use caching::{
Cache, CacheBackend, CacheControls, CacheEntry, CacheKeyContext, CacheKeyField, CacheKeyInput,
CacheMode, cache_key, get_cache, get_cache_key, set_cache, should_use_cache,
pub use cache_type::CacheType;
pub use caching::{Cache, CacheBackend, get_cache, set_cache};
pub use capabilities::{
BatchCache, CacheScript, ClaimCache, ClientInfoCache, CounterCache, DeleteCache, FlushCache,
IncrementOperation, QueueCache, ScanCache, ScriptCache, SetCache, TtlCache,
};
pub use codec::{CacheCodec, JsonCodec};
pub use dual::{DualCache, ReadPolicy, RemoteFailurePolicy, WritePolicy};
pub use error::Error;

View file

@ -1,42 +1,97 @@
use std::{sync::Mutex, time::Duration};
use litellm_cache::{
BaseCache, CacheConnectionResult, CacheControls, CacheEntry, CacheFuture, CacheKeyContext,
CacheKeyField, CacheKeyInput, CacheKwargs, Error, cache_key, get_cache_key,
BaseCache, CacheConnectionResult, CacheContext, Error, ExactCacheContext, get_cache,
};
use sha2::{Digest, Sha256};
use std::time::Duration;
struct TestCache {
default_ttl: Duration,
writes: Mutex<Vec<(String, String, ExactCacheContext)>>,
}
#[derive(Clone)]
struct SemanticContext {
ttl: Option<Duration>,
query: String,
}
impl CacheContext for SemanticContext {
fn ttl(&self) -> Option<Duration> {
self.ttl
}
fn with_ttl(&self, ttl: Option<Duration>) -> Self {
Self {
ttl,
query: self.query.clone(),
}
}
}
struct SemanticCache;
impl BaseCache for SemanticCache {
type Value = String;
type Context = SemanticContext;
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
context.ttl
}
fn set_cache(&self, _: &str, _: Self::Value, _: &Self::Context) -> Result<(), Error> {
Ok(())
}
fn get_cache(&self, _: &str, context: &Self::Context) -> Result<Option<Self::Value>, Error> {
Ok((context.query == "matching prompt").then(|| "semantic hit".into()))
}
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
unreachable!()
}
}
impl BaseCache for TestCache {
type Value = CacheEntry;
type Value = String;
type Context = ExactCacheContext;
fn default_ttl(&self) -> Duration {
self.default_ttl
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
context.ttl.or(Some(self.default_ttl))
}
fn set_cache(&self, _: &str, _: Self::Value, _: CacheKwargs) -> Result<(), Error> {
fn set_cache(&self, _: &str, _: Self::Value, _: &ExactCacheContext) -> Result<(), Error> {
Err(Error::Unavailable)
}
async fn async_set_cache(
&self,
key: &str,
value: Self::Value,
context: ExactCacheContext,
) -> Result<(), Error> {
if key == "unavailable" {
return Err(Error::Unavailable);
}
self.writes
.lock()
.unwrap()
.push((key.into(), value, context));
Ok(())
}
fn get_cache(&self, _: &str, _: &CacheKwargs) -> Result<Option<Self::Value>, Error> {
fn get_cache(&self, _: &str, _: &ExactCacheContext) -> Result<Option<Self::Value>, Error> {
Ok(None)
}
fn delete_cache(&self, _: &str) -> Result<(), Error> {
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
fn flush_cache(&self) -> Result<(), Error> {
Ok(())
}
fn disconnect(&self) -> CacheFuture<'_, ()> {
Box::pin(async { Ok(()) })
}
fn test_connection(&self) -> CacheFuture<'_, CacheConnectionResult> {
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
unreachable!()
}
}
@ -45,95 +100,64 @@ impl BaseCache for TestCache {
fn ttl_uses_default_and_allows_per_call_override() {
let cache = TestCache {
default_ttl: Duration::from_secs(60),
writes: Mutex::default(),
};
assert_eq!(
cache.get_ttl(&CacheKwargs::default()),
Duration::from_secs(60)
cache.get_ttl(&ExactCacheContext::default()),
Some(Duration::from_secs(60))
);
assert_eq!(
cache.get_ttl(&CacheKwargs {
cache.get_ttl(&ExactCacheContext {
ttl: Some(Duration::from_secs(5)),
..Default::default()
}),
Duration::from_secs(5)
Some(Duration::from_secs(5))
);
}
#[test]
fn keys_match_python_order_groups_files_presets_and_namespaces() {
let mut input = CacheKeyInput {
fields: vec![
CacheKeyField {
name: "model".into(),
value: Some("deployment".into()),
api_parameter: true,
internal_parameter: false,
},
CacheKeyField {
name: "file".into(),
value: None,
api_parameter: true,
internal_parameter: false,
},
],
namespace: Some("team".into()),
..Default::default()
fn associated_context_preserves_backend_specific_lookup_inputs() {
let context = SemanticContext {
ttl: None,
query: "matching prompt".into(),
};
CacheKeyContext {
model_group: Some("group".into()),
caching_groups: vec![(vec!["group".into()], "['group']".into())],
file_checksum: Some("checksum".into()),
..Default::default()
}
.apply(&mut input);
assert_eq!(
cache_key(&input),
format!(
"team:{:x}",
Sha256::digest(b"model: ['group']file: checksum")
)
get_cache(&SemanticCache, "shared-key", &context).unwrap(),
Some("semantic hit".into())
);
input.preset = Some("preset".into());
assert_eq!(get_cache_key(&input), "preset");
}
#[test]
fn cache_controls_honor_default_modes_and_directives() {
let enabled = CacheControls {
supported_call_type: true,
configured: true,
default_on: true,
..Default::default()
#[tokio::test]
async fn default_batch_operations_use_async_writes_and_stop_on_failure() {
let cache = TestCache {
default_ttl: Duration::from_secs(60),
writes: Mutex::default(),
};
assert!(enabled.reads());
assert!(enabled.writes());
assert!(
!CacheControls {
default_on: false,
..enabled
}
.reads()
let entry = String::from("cached");
let context = ExactCacheContext {
ttl: Some(Duration::from_secs(5)),
};
cache
.batch_cache_write("single", entry.clone(), context.clone())
.await
.unwrap();
assert_eq!(
cache
.async_set_cache_pipeline(
vec![
("first".into(), entry.clone()),
("unavailable".into(), entry.clone()),
("skipped".into(), entry.clone()),
],
context.clone(),
)
.await,
Err(Error::Unavailable)
);
assert!(
CacheControls {
default_on: false,
use_cache: true,
..enabled
}
.reads()
);
assert!(
!CacheControls {
no_cache: true,
..enabled
}
.reads()
);
assert!(
!CacheControls {
no_store: true,
..enabled
}
.writes()
assert_eq!(
*cache.writes.lock().unwrap(),
vec![
("single".into(), entry.clone(), context.clone()),
("first".into(), entry, context),
]
);
}

View file

@ -0,0 +1,41 @@
use std::collections::BTreeMap;
use litellm_cache::{CacheCodec, Error, JsonCodec};
use serde::{Deserialize, Serialize};
use serde_json::json;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
struct RoutingState {
deployment: String,
cooldown_seconds: u64,
}
#[test]
fn json_codec_round_trips_typed_domain_values() {
let codec = JsonCodec::<RoutingState>::new();
let value = RoutingState {
deployment: "deployment-a".into(),
cooldown_seconds: 30,
};
let bytes = codec.encode(&value).unwrap();
assert_eq!(codec.decode(&bytes).unwrap(), value);
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&bytes).unwrap(),
json!({"deployment": "deployment-a", "cooldown_seconds": 30})
);
}
#[test]
fn json_codec_rejects_malformed_and_wrongly_typed_entries() {
let codec = JsonCodec::<RoutingState>::new();
for bytes in [b"not json".as_slice(), br#"{"deployment":12}"#.as_slice()] {
assert_eq!(codec.decode(bytes).unwrap_err(), Error::InvalidEntry);
}
}
#[test]
fn json_codec_propagates_encoding_errors() {
let codec = JsonCodec::<BTreeMap<(u8, u8), String>>::new();
let value = BTreeMap::from([((1, 2), "invalid JSON object key".into())]);
assert_eq!(codec.encode(&value).unwrap_err(), Error::InvalidEntry);
}

385
litellm-rust/crates/cache/tests/dual.rs vendored Normal file
View file

@ -0,0 +1,385 @@
use std::{
sync::{Arc, Mutex},
time::Duration,
};
use litellm_cache::{
BaseCache, BatchCache, CacheConnectionResult, ClaimCache, CounterCache, DeleteCache, DualCache,
Error, ExactCacheContext, FlushCache, ReadPolicy, RemoteFailurePolicy, WritePolicy,
};
struct TestCache<V> {
value: Mutex<Option<V>>,
fail: bool,
}
impl<V> TestCache<V> {
fn new(value: Option<V>, fail: bool) -> Self {
Self {
value: Mutex::new(value),
fail,
}
}
}
impl<V> BaseCache for TestCache<V>
where
V: Clone + Send + Sync + 'static,
{
type Value = V;
type Context = ExactCacheContext;
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
context.ttl.or(Some(Duration::from_secs(60)))
}
fn set_cache(&self, _: &str, value: V, _: &ExactCacheContext) -> Result<(), Error> {
*self.value.lock().unwrap() = Some(value);
Ok(())
}
fn get_cache(&self, _: &str, _: &ExactCacheContext) -> Result<Option<V>, Error> {
Ok(self.value.lock().unwrap().clone())
}
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
unreachable!()
}
}
impl<V> BatchCache for TestCache<V> where V: Clone + Send + Sync + 'static {}
impl<V> DeleteCache for TestCache<V>
where
V: Clone + Send + Sync + 'static,
{
fn delete_cache(&self, _: &str) -> Result<(), Error> {
*self.value.lock().unwrap() = None;
Ok(())
}
}
impl<V> FlushCache for TestCache<V>
where
V: Clone + Send + Sync + 'static,
{
fn flush_cache(&self) -> Result<(), Error> {
*self.value.lock().unwrap() = None;
Ok(())
}
}
impl CounterCache for TestCache<f64> {
fn increment_cache(&self, _: &str, amount: f64, _: ExactCacheContext) -> Result<f64, Error> {
if self.fail {
return Err(Error::Unavailable);
}
let mut value = self.value.lock().unwrap();
let incremented = value.unwrap_or_default() + amount;
*value = Some(incremented);
Ok(incremented)
}
}
impl<V> ClaimCache for TestCache<V>
where
V: Clone + PartialEq + Send + Sync + 'static,
{
fn claim_cache(
&self,
_: &str,
candidate: V,
eligible: &[V],
_: ExactCacheContext,
) -> Result<V, Error> {
if self.fail {
return Err(Error::Unavailable);
}
let mut value = self.value.lock().unwrap();
let winner = match value.as_ref() {
Some(existing) if eligible.is_empty() || eligible.contains(existing) => {
existing.clone()
}
_ => candidate,
};
*value = Some(winner.clone());
Ok(winner)
}
}
#[test]
fn failed_l2_increment_leaves_l1_unchanged() {
let l1 = Arc::new(TestCache::new(Some(10.0), false));
let cache = DualCache::new(l1.clone(), Arc::new(TestCache::new(Some(20.0), true)));
assert_eq!(
cache.increment_cache("counter", 2.0, ExactCacheContext::default()),
Err(Error::Unavailable)
);
assert_eq!(
l1.get_cache("counter", &ExactCacheContext::default())
.unwrap(),
Some(10.0)
);
}
#[test]
fn claim_uses_l1_fallback_without_overwriting_an_eligible_winner() {
let l1 = Arc::new(TestCache::new(Some("first".to_string()), false));
let cache = DualCache::new(l1, Arc::new(TestCache::new(None, true)))
.with_remote_failure_policy(RemoteFailurePolicy::UseLocal);
assert_eq!(
cache
.claim_cache(
"affinity",
"second".into(),
&["first".into(), "second".into()],
ExactCacheContext {
ttl: Some(Duration::from_secs(60)),
},
)
.unwrap(),
"first"
);
}
struct SyncPanics(TestCache<String>);
impl BaseCache for SyncPanics {
type Value = String;
type Context = ExactCacheContext;
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
self.0.get_ttl(context)
}
fn set_cache(&self, _: &str, _: String, _: &ExactCacheContext) -> Result<(), Error> {
panic!("sync L2 write on an async path")
}
fn get_cache(&self, _: &str, _: &ExactCacheContext) -> Result<Option<String>, Error> {
panic!("sync L2 read on an async path")
}
async fn async_set_cache(
&self,
key: &str,
value: String,
context: ExactCacheContext,
) -> Result<(), Error> {
self.0.set_cache(key, value, &context)
}
async fn async_get_cache(
&self,
key: &str,
context: &ExactCacheContext,
) -> Result<Option<String>, Error> {
self.0.get_cache(key, context)
}
async fn async_set_cache_pipeline(
&self,
cache_list: Vec<(String, String)>,
context: ExactCacheContext,
) -> Result<(), Error> {
for (key, value) in cache_list {
self.0.set_cache(&key, value, &context)?;
}
Ok(())
}
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
unreachable!()
}
}
impl BatchCache for SyncPanics {
async fn async_batch_get_cache(
&self,
keys: Vec<String>,
context: ExactCacheContext,
) -> Result<Vec<litellm_cache::BatchEntry<String>>, Error> {
assert_eq!(keys, ["missing"]);
Ok(vec![match self.0.get_cache("missing", &context)? {
Some(value) => litellm_cache::BatchEntry::Hit(value),
None => litellm_cache::BatchEntry::Miss,
}])
}
}
impl DeleteCache for SyncPanics {
fn delete_cache(&self, _: &str) -> Result<(), Error> {
panic!("sync L2 delete on an async path")
}
async fn async_delete_cache(&self, key: &str) -> Result<(), Error> {
self.0.delete_cache(key)
}
}
impl FlushCache for SyncPanics {
fn flush_cache(&self) -> Result<(), Error> {
panic!("sync L2 flush on an async path")
}
}
#[tokio::test]
async fn async_operations_use_the_async_l2_methods() {
let l1 = Arc::new(TestCache::new(None, false));
let cache = DualCache::new(
l1.clone(),
Arc::new(SyncPanics(TestCache::new(
Some("remote".to_string()),
false,
))),
);
let context = ExactCacheContext::default();
assert_eq!(
cache.async_get_cache("missing", &context).await.unwrap(),
Some("remote".into())
);
assert_eq!(
l1.get_cache("missing", &context).unwrap(),
Some("remote".into())
);
l1.delete_cache("missing").unwrap();
assert_eq!(
cache
.async_batch_get_cache(vec!["missing".into()], context.clone())
.await
.unwrap(),
[litellm_cache::BatchEntry::Hit("remote".to_string())]
);
cache
.async_set_cache("missing", "written".into(), context.clone())
.await
.unwrap();
cache
.async_set_cache_pipeline(vec![("missing".into(), "piped".into())], context.clone())
.await
.unwrap();
cache.async_delete_cache("missing").await.unwrap();
assert_eq!(
cache.async_get_cache("missing", &context).await.unwrap(),
None
);
}
struct Unavailable;
impl BaseCache for Unavailable {
type Value = String;
type Context = ExactCacheContext;
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
context.ttl
}
fn set_cache(&self, _: &str, _: String, _: &ExactCacheContext) -> Result<(), Error> {
Err(Error::Unavailable)
}
fn get_cache(&self, _: &str, _: &ExactCacheContext) -> Result<Option<String>, Error> {
Err(Error::Unavailable)
}
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
unreachable!()
}
}
impl BatchCache for Unavailable {}
impl DeleteCache for Unavailable {
fn delete_cache(&self, _: &str) -> Result<(), Error> {
Err(Error::Unavailable)
}
}
impl FlushCache for Unavailable {
fn flush_cache(&self) -> Result<(), Error> {
Err(Error::Unavailable)
}
}
impl ClaimCache for Unavailable {
fn claim_cache(
&self,
_: &str,
_: String,
_: &[String],
_: ExactCacheContext,
) -> Result<String, Error> {
Err(Error::InvalidEntry)
}
}
#[test]
fn remote_failure_policy_selects_propagation_or_the_local_tier() {
let context = ExactCacheContext::default();
let strict = DualCache::new(Arc::new(TestCache::new(None, false)), Arc::new(Unavailable));
assert_eq!(
strict.set_cache("key", "value".into(), &context),
Err(Error::Unavailable)
);
assert_eq!(strict.get_cache("key", &context), Err(Error::Unavailable));
let l1 = Arc::new(TestCache::new(None, false));
let degraded = DualCache::new(l1.clone(), Arc::new(Unavailable))
.with_remote_failure_policy(RemoteFailurePolicy::UseLocal);
assert_eq!(degraded.get_cache("key", &context), Ok(None));
degraded.set_cache("key", "value".into(), &context).unwrap();
assert_eq!(
degraded.get_cache("key", &context),
Ok(Some("value".into()))
);
degraded.delete_cache("key").unwrap();
assert_eq!(l1.get_cache("key", &context), Ok(None));
}
#[test]
fn claim_fallback_does_not_hide_non_availability_errors() {
let cache = DualCache::new(
Arc::new(TestCache::new(Some("first".to_string()), false)),
Arc::new(Unavailable),
)
.with_remote_failure_policy(RemoteFailurePolicy::UseLocal);
assert_eq!(
cache.claim_cache(
"affinity",
"second".into(),
&[],
ExactCacheContext::default()
),
Err(Error::InvalidEntry)
);
}
#[test]
fn local_only_policies_never_touch_l2() {
let l2 = Arc::new(TestCache::new(Some("remote".to_string()), false));
let cache = DualCache::new(Arc::new(TestCache::new(None, false)), l2.clone())
.with_read_policy(ReadPolicy::LocalOnly)
.with_write_policy(WritePolicy::LocalOnly);
let context = ExactCacheContext::default();
assert_eq!(cache.get_cache("key", &context), Ok(None));
cache.set_cache("key", "local".into(), &context).unwrap();
assert_eq!(l2.get_cache("key", &context), Ok(Some("remote".into())));
}

View file

@ -1,33 +1,91 @@
use serde::{Deserialize, Deserializer, de::Error};
use serde_json::Value;
use serde::{
Deserializer,
de::{Error, Visitor},
};
use serde_with::DeserializeAs;
pub struct LaxI64;
pub struct FiniteF64;
pub fn parse_str_bool(value: &str) -> Option<bool> {
let token = value.trim_matches(|character: char| {
character.is_whitespace() || matches!(character, '\u{1c}'..='\u{1f}')
});
if token.eq_ignore_ascii_case("true") {
return Some(true);
}
token.eq_ignore_ascii_case("false").then_some(false)
}
impl<'de> DeserializeAs<'de, i64> for LaxI64 {
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<i64, D::Error> {
match Value::deserialize(deserializer)? {
Value::Number(number) if number.is_f64() => number.as_f64().and_then(integral_float),
Value::Number(number) => number.as_i64(),
Value::String(value) => integer_string(value.trim()),
Value::Bool(value) => Some(i64::from(value)),
_ => None,
}
.ok_or_else(|| D::Error::custom("expected an integer in the i64 range"))
deserializer.deserialize_any(Self)
}
}
impl<'de> Visitor<'de> for LaxI64 {
type Value = i64;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("an integer in the i64 range")
}
fn visit_i64<E: Error>(self, value: i64) -> Result<i64, E> {
Ok(value)
}
fn visit_u64<E: Error>(self, value: u64) -> Result<i64, E> {
i64::try_from(value).map_err(E::custom)
}
fn visit_f64<E: Error>(self, value: f64) -> Result<i64, E> {
integral_float(value).ok_or_else(|| E::custom("expected an integer in the i64 range"))
}
fn visit_str<E: Error>(self, value: &str) -> Result<i64, E> {
integer_string(value.trim())
.ok_or_else(|| E::custom("expected an integer in the i64 range"))
}
fn visit_bool<E: Error>(self, value: bool) -> Result<i64, E> {
Ok(i64::from(value))
}
}
impl<'de> DeserializeAs<'de, f64> for FiniteF64 {
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<f64, D::Error> {
match Value::deserialize(deserializer)? {
Value::Number(number) => number.as_f64(),
Value::String(value) => value.trim().parse::<f64>().ok(),
Value::Bool(value) => Some(f64::from(value)),
_ => None,
}
.filter(|value| value.is_finite())
.ok_or_else(|| D::Error::custom("expected a finite number"))
deserializer.deserialize_any(Self)
}
}
impl<'de> Visitor<'de> for FiniteF64 {
type Value = f64;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a finite number")
}
fn visit_i64<E: Error>(self, value: i64) -> Result<f64, E> {
Ok(value as f64)
}
fn visit_u64<E: Error>(self, value: u64) -> Result<f64, E> {
Ok(value as f64)
}
fn visit_f64<E: Error>(self, value: f64) -> Result<f64, E> {
value
.is_finite()
.then_some(value)
.ok_or_else(|| E::custom("expected a finite number"))
}
fn visit_str<E: Error>(self, value: &str) -> Result<f64, E> {
self.visit_f64(value.trim().parse::<f64>().map_err(E::custom)?)
}
fn visit_bool<E: Error>(self, value: bool) -> Result<f64, E> {
Ok(f64::from(value))
}
}
@ -66,7 +124,7 @@ fn integral_float(value: f64) -> Option<i64> {
#[cfg(test)]
mod tests {
use serde::Serialize;
use serde::{Deserialize, Serialize};
use serde_json::json;
use serde_with::serde_as;
@ -81,6 +139,22 @@ mod tests {
float: Option<f64>,
}
#[test]
fn boolean_tokens_follow_python_string_trimming_without_redis_tokens() {
for (input, expected) in [
(" True ", Some(true)),
("\u{1c}TRUE\u{1f}", Some(true)),
("\u{a0}False\u{2003}", Some(false)),
("true\u{200b}", None),
("yes", None),
("1", None),
("", None),
("unknown", None),
] {
assert_eq!(parse_str_bool(input), expected, "{input:?}");
}
}
#[test]
fn adapters_compose_and_serialize_as_numbers() {
let numbers: Numbers = serde_json::from_value(json!({

View file

@ -1,5 +1,7 @@
use std::str::FromStr;
use crate::serde_compat::parse_str_bool;
pub trait Lookup {
fn get(&self, name: &str) -> Option<String>;
@ -9,7 +11,7 @@ pub trait Lookup {
fn enabled(&self, name: &str) -> Option<bool> {
self.get(name)
.is_some_and(|value| value.trim().eq_ignore_ascii_case("true"))
.is_some_and(|value| parse_str_bool(&value) == Some(true))
.then_some(true)
}

View file

@ -45,7 +45,7 @@ where
fn into_pyobject(self, py: Python<'py>) -> PyResult<Self::Output> {
catch_unwind(AssertUnwindSafe(|| pythonize::pythonize(py, &self.0)))
.map_err(panic_to_pyerr)?
.map_err(|error| PyValueError::new_err(error.to_string()))
.map_err(PyErr::from)
}
}
@ -87,6 +87,19 @@ mod tests {
});
}
#[test]
fn pythonized_preserves_python_serialization_error_types() {
crate::initialize_python();
Python::attach(|py| {
let value = std::collections::BTreeMap::from([(vec![1], "value")]);
let direct = to_py(py, &value).unwrap_err();
let wrapped = Pythonized(value).into_pyobject(py).unwrap_err();
assert!(direct.is_instance_of::<pyo3::exceptions::PyTypeError>(py));
assert!(wrapped.is_instance_of::<pyo3::exceptions::PyTypeError>(py));
assert_eq!(wrapped.to_string(), direct.to_string());
});
}
#[test]
fn pythonized_maps_serializer_panics_to_a_base_exception() {
crate::initialize_python();

View file

@ -129,6 +129,7 @@ mod tests {
use rstest::rstest;
use super::*;
use crate::TlsSource;
fn settings(ssl_verify: Option<SslVerify>, ssl_cert_file: Option<&str>) -> HttpSettings {
HttpSettings {
@ -298,7 +299,11 @@ mod tests {
};
assert!(matches!(
reqwest::ClientBuilder::try_from(&config),
Err(Error::Read { path: reported, .. }) if reported == path
Err(Error::Read {
path: reported,
tls_source: TlsSource::CaBundle,
..
}) if reported == path
));
}
@ -315,7 +320,11 @@ mod tests {
std::fs::remove_file(&path).unwrap();
assert!(matches!(
result,
Err(Error::InvalidPem { path: reported, .. }) if reported == path
Err(Error::InvalidPem {
path: reported,
tls_source: TlsSource::CaBundle,
..
}) if reported == path
));
}
}

View file

@ -1,11 +1,25 @@
use std::path::PathBuf;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum TlsSource {
CaBundle,
ClientIdentity,
}
#[derive(Clone, Debug, thiserror::Error, PartialEq, Eq)]
pub enum Error {
#[error("could not read {}: {message}", path.display())]
Read { path: PathBuf, message: String },
Read {
path: PathBuf,
message: String,
tls_source: TlsSource,
},
#[error("{} is not a PEM file: {message}", path.display())]
InvalidPem { path: PathBuf, message: String },
InvalidPem {
path: PathBuf,
message: String,
tls_source: TlsSource,
},
#[error("could not build the HTTP client: {0}")]
Client(String),
#[error("request body could not be serialized: {0}")]

View file

@ -10,7 +10,7 @@ mod tls;
pub mod transport;
pub use config::{HttpClientConfig, Resolution, Verify};
pub use error::Error;
pub use error::{Error, TlsSource};
pub use pool::{ClientVariant, HttpClientPool};
pub use proxy::EnvironmentProxies;
pub use settings::{HttpSettings, HttpSettingsLayer, SslVerify, TcpKeepalive};

View file

@ -54,16 +54,39 @@ impl Default for UrlPolicy {
impl UrlPolicy {
fn allows(&self, host: &str, port: u16) -> bool {
let host = normalize_host(host);
let with_port = format!("{host}:{port}");
self.allowed_hosts
.iter()
.map(|entry| normalize_host(entry))
.any(|entry| entry == host || entry == with_port)
.filter_map(|entry| parse_allowed_host(entry))
.any(|(entry_host, entry_port)| {
entry_host == host && entry_port.is_none_or(|entry_port| entry_port == port)
})
}
}
fn normalize_host(host: &str) -> String {
host.to_ascii_lowercase().trim_end_matches('.').to_owned()
pub fn normalize_host(host: &str) -> String {
let host = host.trim().trim_end_matches('.');
let host = host
.strip_prefix('[')
.and_then(|host| host.strip_suffix(']'))
.unwrap_or(host);
host.to_ascii_lowercase()
}
fn parse_allowed_host(entry: &str) -> Option<(String, Option<u16>)> {
let entry = entry.trim();
if let Some(entry) = entry.strip_prefix('[') {
let (host, suffix) = entry.split_once(']')?;
let port = match suffix {
"" => None,
suffix => Some(suffix.strip_prefix(':')?.parse().ok()?),
};
return Some((normalize_host(host), port));
}
let (host, port) = match entry.rsplit_once(':') {
Some((host, port)) if !host.contains(':') => (host, Some(port.parse().ok()?)),
_ => (entry, None),
};
Some((normalize_host(host), port))
}
type ProxyMatch = Arc<dyn Fn(&Url) -> bool + Send + Sync>;
@ -670,6 +693,21 @@ mod tests {
assert!(matches!(result, Err(Error::BlockedUrl)));
}
#[test]
fn allowlist_matches_bracketed_ipv6_hosts_and_ports() {
let policy = UrlPolicy {
validate: true,
allowed_hosts: vec!["[2001:db8::1]".into(), "[2001:db8::1]:8443".into()],
};
assert!(policy.allows("2001:db8::1", 443));
assert!(policy.allows("2001:db8::1", 8443));
let port_specific = UrlPolicy {
validate: true,
allowed_hosts: vec!["[2001:db8::1]:8443".into()],
};
assert!(!port_specific.allows("2001:db8::1", 9443));
}
#[tokio::test]
async fn validation_off_fetches_private_hosts_and_follows_redirects() {
let (url, server, _) = serve_named(

View file

@ -3,7 +3,10 @@ use std::{
time::Duration,
};
use litellm_core_utils::settings::{Layer, Lookup, merge};
use litellm_core_utils::{
serde_compat::parse_str_bool,
settings::{Layer, Lookup, merge},
};
use crate::proxy::EnvironmentProxies;
@ -16,9 +19,9 @@ pub enum SslVerify {
impl SslVerify {
pub fn parse(value: &str) -> Self {
match value.trim().to_ascii_lowercase().as_str() {
"true" => Self::Enabled,
"false" => Self::Disabled,
match parse_str_bool(value) {
Some(true) => Self::Enabled,
Some(false) => Self::Disabled,
_ => Self::CaBundle(PathBuf::from(value)),
}
}
@ -152,9 +155,7 @@ impl HttpSettings {
Self {
ssl_verify: merged.ssl_verify,
ssl_cert_file: merged.ssl_cert_file,
ssl_certificate: merged
.ssl_certificate
.filter(|path| !path.as_os_str().is_empty()),
ssl_certificate: merged.ssl_certificate,
ssl_security_level: merged.ssl_security_level.filter(|level| !level.is_empty()),
ssl_ecdh_curve: merged.ssl_ecdh_curve.filter(|curve| !curve.is_empty()),
force_ipv4: merged.force_ipv4.unwrap_or(defaults.force_ipv4),
@ -287,7 +288,7 @@ mod tests {
}
#[test]
fn empty_environment_values_clear_the_setting_like_python_truthiness() {
fn empty_certificate_is_retained_for_validation_while_empty_tuning_is_absent() {
let configured = HttpSettingsLayer {
ssl_certificate: Some("/configured/client.pem".into()),
ssl_security_level: Some("configured".into()),
@ -300,7 +301,7 @@ mod tests {
("SSL_ECDH_CURVE", ""),
]));
let settings = HttpSettings::from_layers([environment, configured]);
assert_eq!(settings.ssl_certificate, None);
assert_eq!(settings.ssl_certificate, Some(PathBuf::new()));
assert_eq!(settings.ssl_security_level, None);
assert_eq!(settings.ssl_ecdh_curve, None);
}

View file

@ -9,7 +9,7 @@ use rustls::{
use crate::{
config::{HttpClientConfig, Verify},
error::Error,
error::{Error, TlsSource},
};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
@ -197,15 +197,17 @@ impl TryFrom<&HttpClientConfig> for ClientConfig {
Verify::BuiltInRoots => builder.with_root_certificates(RootCertStore {
roots: webpki_roots::TLS_SERVER_ROOTS.to_vec(),
}),
Verify::CaBundle(path) => builder.with_root_certificates(bundle_roots(path)?),
Verify::CaBundle(path) => {
builder.with_root_certificates(bundle_roots(path, TlsSource::CaBundle)?)
}
};
let mut tls = match &config.client_certificate {
None => verified.with_no_client_auth(),
Some(path) => {
let (chain, key) = identity(path)?;
let (chain, key) = identity(path, TlsSource::ClientIdentity)?;
verified
.with_client_auth_cert(chain, key)
.map_err(|error| invalid_pem(path, error))?
.map_err(|error| invalid_pem(path, TlsSource::ClientIdentity, error))?
}
};
tls.alpn_protocols = if config.http2 {
@ -217,47 +219,52 @@ impl TryFrom<&HttpClientConfig> for ClientConfig {
}
}
fn bundle_roots(path: &Path) -> Result<RootCertStore, Error> {
let certificates = certificates(path)?;
fn bundle_roots(path: &Path, source: TlsSource) -> Result<RootCertStore, Error> {
let certificates = certificates(path, source)?;
if certificates.is_empty() {
return Err(invalid_pem(path, "no certificates found"));
return Err(invalid_pem(path, source, "no certificates found"));
}
let mut store = RootCertStore::empty();
for certificate in certificates {
store
.add(certificate)
.map_err(|error| invalid_pem(path, error))?;
.map_err(|error| invalid_pem(path, source, error))?;
}
Ok(store)
}
fn identity(path: &Path) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>), Error> {
let chain = certificates(path)?;
fn identity(
path: &Path,
source: TlsSource,
) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>), Error> {
let chain = certificates(path, source)?;
if chain.is_empty() {
return Err(invalid_pem(path, "no certificates found"));
return Err(invalid_pem(path, source, "no certificates found"));
}
let key =
PrivateKeyDer::from_pem_slice(&read(path)?).map_err(|error| invalid_pem(path, error))?;
let key = PrivateKeyDer::from_pem_slice(&read(path, source)?)
.map_err(|error| invalid_pem(path, source, error))?;
Ok((chain, key))
}
fn certificates(path: &Path) -> Result<Vec<CertificateDer<'static>>, Error> {
CertificateDer::pem_slice_iter(&read(path)?)
fn certificates(path: &Path, source: TlsSource) -> Result<Vec<CertificateDer<'static>>, Error> {
CertificateDer::pem_slice_iter(&read(path, source)?)
.collect::<Result<_, _>>()
.map_err(|error| invalid_pem(path, error))
.map_err(|error| invalid_pem(path, source, error))
}
fn read(path: &Path) -> Result<Vec<u8>, Error> {
fn read(path: &Path, source: TlsSource) -> Result<Vec<u8>, Error> {
std::fs::read(path).map_err(|error| Error::Read {
path: path.to_path_buf(),
message: error.to_string(),
tls_source: source,
})
}
fn invalid_pem(path: &Path, message: impl fmt::Display) -> Error {
fn invalid_pem(path: &Path, source: TlsSource, message: impl fmt::Display) -> Error {
Error::InvalidPem {
path: path.to_path_buf(),
message: message.to_string(),
tls_source: source,
}
}
@ -405,7 +412,11 @@ mod tests {
std::fs::remove_file(&path).unwrap();
assert!(matches!(
result,
Err(Error::InvalidPem { path: reported, .. }) if reported == path
Err(Error::InvalidPem {
path: reported,
tls_source: TlsSource::ClientIdentity,
..
}) if reported == path
));
}
}

View file

@ -20,6 +20,11 @@ tiktoken = ["litellm-token-counter/tiktoken"]
[dependencies]
bytes.workspace = true
litellm-cache.workspace = true
litellm-cache-memory.workspace = true
litellm-cache-redis.workspace = true
litellm-cache-response.workspace = true
serde.workspace = true
litellm-auth.workspace = true
litellm-callbacks-legacy-python.workspace = true
litellm-core.workspace = true
@ -36,6 +41,8 @@ serde_json.workspace = true
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
serde.workspace = true
serde_with.workspace = true
criterion.workspace = true
futures-util.workspace = true
rstest.workspace = true

View file

@ -1,26 +1,154 @@
{
"http_settings": [
"ssl_verify",
"ssl_certificate",
"ssl_security_level",
"ssl_ecdh_curve",
"force_ipv4",
"http2",
"aiohttp_trust_env",
"disable_aiohttp_trust_env",
"disable_aiohttp_transport",
"user_agent"
],
"url_policy": [
"user_url_validation",
"user_url_allowed_hosts"
],
"provider_defaults": [
"vertex_project",
"vertex_location",
"enable_azure_ad_token_refresh"
],
"secret_manager": [
"readable"
]
"http_settings": {
"version": 1,
"fields": {
"ssl_verify": {
"adapter": "SslVerifyInput",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [
"none",
"bool",
"str"
],
"unsupported_live": "configuration_error"
},
"ssl_certificate": {
"adapter": "OptionalStrictString",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"ssl_security_level": {
"adapter": "TuningString",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"ssl_ecdh_curve": {
"adapter": "TuningString",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"force_ipv4": {
"adapter": "Truthy",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"http2": {
"adapter": "ExactTrue",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"aiohttp_trust_env": {
"adapter": "Truthy",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"disable_aiohttp_trust_env": {
"adapter": "Truthy",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"disable_aiohttp_transport": {
"adapter": "ExactTrue",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"user_agent": {
"adapter": "StrictString",
"required": true,
"precedence": "accessor",
"sensitive": false,
"shapes": [],
"unsupported_live": null
}
}
},
"url_policy": {
"version": 1,
"fields": {
"user_url_validation": {
"adapter": "Truthy",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
},
"user_url_allowed_hosts": {
"adapter": "HostCollection",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
}
}
},
"provider_defaults": {
"version": 1,
"fields": {
"vertex_project": {
"adapter": "FalsyOptionalString",
"required": true,
"precedence": "module_global",
"sensitive": true,
"shapes": [],
"unsupported_live": null
},
"vertex_location": {
"adapter": "FalsyOptionalString",
"required": true,
"precedence": "module_global",
"sensitive": true,
"shapes": [],
"unsupported_live": null
},
"enable_azure_ad_token_refresh": {
"adapter": "ExactTrue",
"required": true,
"precedence": "module_global",
"sensitive": false,
"shapes": [],
"unsupported_live": null
}
}
},
"secret_manager": {
"version": 1,
"fields": {
"readable": {
"adapter": "StrictBool",
"required": true,
"precedence": "accessor",
"sensitive": false,
"shapes": [],
"unsupported_live": null
}
}
}
}

View file

@ -0,0 +1,291 @@
use litellm_cache_response::PartialHits;
use litellm_host_python::{ExecutionStep, from_py, release_gil, run_async, to_py};
use pyo3::{
PyTraverseError, PyVisit,
exceptions::{PyRuntimeError, PyValueError},
prelude::*,
types::PyDict,
};
use serde_json::Value;
use super::{
cache_error,
callback::PythonCallback,
future::{ready_none, ready_value},
native::NativeResponseCache,
request::{now, request, requests},
};
pub(super) enum CacheBinding {
Disabled,
Native(NativeResponseCache),
PythonCallback(PythonCallback),
}
#[pyclass(frozen, name = "_CacheTestBinding")]
pub(crate) struct ResolvedCache {
binding: CacheBinding,
pid: u32,
}
impl ResolvedCache {
pub(super) fn new(binding: CacheBinding) -> Self {
Self {
binding,
pid: std::process::id(),
}
}
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)?;
let service = service.clone();
run_async(
py,
async move { service.async_lookup(&request, now()).await },
cache_error,
)?
}
CacheBinding::PythonCallback(callback) => callback.async_lookup(py, kwargs)?,
};
Ok(ExecutionStep::Await(awaitable.unbind()))
}
}
#[pymethods]
impl ResolvedCache {
#[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)
}
}
}
#[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)?;
let service = service.clone();
run_async(
py,
async move { service.async_store(&request, response, now()).await },
cache_error,
)
}
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();
let service = service.clone();
run_async(
py,
async move { service.async_store_batch(entries, now()).await },
cache_error,
)
}
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)?;
}
Ok(())
}
}

View file

@ -0,0 +1,162 @@
use pyo3::{
PyTraverseError, PyVisit,
exceptions::{PyTypeError, PyValueError},
prelude::*,
types::{PyDict, PyList, PyTuple},
};
use super::future::ready_none;
pub(super) struct PythonCallback(Py<PyAny>);
impl PythonCallback {
pub(super) fn new(object: Py<PyAny>) -> Self {
Self(object)
}
pub(super) 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(super) 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(super) 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(super) 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(super) 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(super) 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(super) 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(super) 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(super) fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
self.0.bind(py).call_method0("ping")
}
pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.0)
}
}
fn callback_kwargs<'a, 'py>(
kwargs: Option<&'a Bound<'py, PyDict>>,
) -> PyResult<&'a Bound<'py, PyDict>> {
kwargs.ok_or_else(|| {
PyTypeError::new_err("Python cache callbacks require their original callback_kwargs")
})
}
fn batch_callback_kwargs<'py>(
requests: &Bound<'py, PyAny>,
kwargs: Option<&Bound<'py, PyAny>>,
) -> PyResult<Vec<Bound<'py, PyDict>>> {
let kwargs = kwargs
.ok_or_else(|| {
PyTypeError::new_err(
"Python cache callbacks require one original callback_kwargs mapping per request",
)
})?
.try_iter()?
.map(|item| Ok(item?.cast_into::<PyDict>()?))
.collect::<PyResult<Vec<_>>>()?;
if kwargs.len() != requests.len()? {
return Err(PyValueError::new_err(
"batch cache requests and callback_kwargs must have equal lengths",
));
}
Ok(kwargs)
}

View file

@ -0,0 +1,826 @@
use std::time::Duration;
use litellm_cache::CacheType;
use litellm_cache_redis::{RedisNode, RedisTopology};
use pyo3::{
exceptions::{PyTypeError, PyValueError},
prelude::*,
types::{PyAny, PyDict, PyList, PyString},
};
use super::{native::NativeResponseCache, request::duration};
#[allow(dead_code, reason = "consumed by the cache activation follow-up")]
pub(super) struct CachePolicy {
pub(super) mode: String,
pub(super) ttl: Option<Duration>,
pub(super) namespace: Option<String>,
pub(super) supported_call_types: Option<Vec<String>>,
pub(super) redis_flush_size: Option<usize>,
pub(super) semantic_cache_scope: String,
}
pub(super) struct MemoryCacheConfig {
pub(super) default_ttl: Duration,
pub(super) capacity: usize,
pub(super) max_entry_bytes: usize,
}
#[derive(Debug, PartialEq)]
pub(super) enum RedisProtocol {
Resp2,
Resp3,
}
#[derive(Debug, PartialEq)]
pub(super) enum CertificateRequirement {
None,
Optional,
Required,
}
#[allow(dead_code, reason = "consumed by the cache activation follow-up")]
pub(super) struct RedisTlsConfig {
pub(super) certificate_requirement: CertificateRequirement,
pub(super) check_hostname: bool,
pub(super) ca_certificate: Option<String>,
pub(super) ca_data: Option<String>,
pub(super) client_certificate: Option<String>,
pub(super) client_key: Option<String>,
}
#[allow(dead_code, reason = "consumed by the cache activation follow-up")]
pub(super) struct RedisConnectionConfig {
pub(super) host: String,
pub(super) port: u16,
pub(super) database: i64,
pub(super) username: Option<String>,
pub(super) password: Option<String>,
pub(super) protocol: RedisProtocol,
pub(super) pool_size: usize,
pub(super) read_timeout: Option<Duration>,
pub(super) connect_timeout: Option<Duration>,
pub(super) socket_keepalive: Option<bool>,
pub(super) health_check_interval: Duration,
pub(super) client_name: Option<String>,
pub(super) tls: Option<RedisTlsConfig>,
}
#[allow(dead_code, reason = "consumed by the cache activation follow-up")]
pub(super) struct RedisCacheConfig {
pub(super) default_ttl: Duration,
pub(super) namespace: Option<String>,
pub(super) flush_size: usize,
pub(super) topology: RedisTopology,
pub(super) connection: RedisConnectionConfig,
}
struct RedisClientProjection<'py> {
topology: RedisTopology,
host: String,
port: u16,
pool_size: usize,
resolved: Bound<'py, PyDict>,
tls: Option<RedisTlsConfig>,
}
const REDIS_PY_DEFAULT_MAX_CONNECTIONS: usize = 1 << 31;
pub(super) enum CacheBackendConfig {
Memory(MemoryCacheConfig),
Redis(Box<RedisCacheConfig>),
}
#[allow(dead_code, reason = "consumed by the cache activation follow-up")]
pub(super) struct NativeCacheConfig {
pub(super) policy: CachePolicy,
pub(super) backend: CacheBackendConfig,
}
pub(super) enum UnsupportedCacheConfig {
Backend,
RedisTopology,
RedisCredentials,
RedisConnection,
RedisOption,
}
impl UnsupportedCacheConfig {
pub(super) fn message(&self) -> &'static str {
match self {
Self::Backend => "native cache backend is not implemented",
Self::RedisTopology => "native Redis topology is not implemented",
Self::RedisCredentials => "native Redis credentials require Python",
Self::RedisConnection => "native Redis connection type is not implemented",
Self::RedisOption => "native Redis configuration requires Python",
}
}
}
pub(super) enum CacheConfigProjection {
Native(Box<NativeCacheConfig>),
Unsupported(UnsupportedCacheConfig),
}
impl NativeCacheConfig {
#[inline(never)]
pub(super) fn project(facade: &Bound<'_, PyAny>) -> PyResult<CacheConfigProjection> {
let backend_name = facade.getattr("type")?.extract::<String>()?;
let policy = CachePolicy {
mode: facade.getattr("mode")?.extract::<String>()?,
ttl: optional_duration(facade.getattr("ttl")?)?,
namespace: optional_string(facade.getattr("namespace")?)?,
supported_call_types: facade
.getattr("supported_call_types")?
.extract::<Option<Vec<String>>>()?,
redis_flush_size: facade
.getattr("redis_flush_size")?
.extract::<Option<usize>>()?,
semantic_cache_scope: facade
.getattr("semantic_cache_scope")?
.extract::<String>()?,
};
let backend = facade.getattr("cache")?;
match CacheType::from_python_name(&backend_name) {
Some(CacheType::Local) => project_memory(&backend).map(|backend| {
CacheConfigProjection::Native(Box::new(Self {
policy,
backend: CacheBackendConfig::Memory(backend),
}))
}),
Some(CacheType::Redis) => match project_redis(&backend)? {
Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self {
policy,
backend: CacheBackendConfig::Redis(Box::new(backend)),
}))),
Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)),
},
Some(
CacheType::RedisSemantic
| CacheType::ValkeySemantic
| CacheType::S3
| CacheType::Disk
| CacheType::QdrantSemantic
| CacheType::AzureBlob
| CacheType::Gcs,
)
| None => Ok(CacheConfigProjection::Unsupported(
UnsupportedCacheConfig::Backend,
)),
}
}
pub(super) fn service_mismatch(&self, service: &NativeResponseCache) -> Option<&'static str> {
if service.default_ttl()
!= Some(match &self.backend {
CacheBackendConfig::Memory(config) => config.default_ttl,
CacheBackendConfig::Redis(config) => config.default_ttl,
})
{
return Some("facade and native backend default TTLs must match");
}
match &self.backend {
CacheBackendConfig::Memory(config) if service.kind() != "memory" => {
Some("facade and native backend types must match")
}
CacheBackendConfig::Memory(config) if service.capacity() != Some(config.capacity) => {
Some("facade and native backend capacities must match")
}
CacheBackendConfig::Memory(config)
if service.max_entry_bytes() != Some(config.max_entry_bytes) =>
{
Some("facade and native backend item limits must match")
}
CacheBackendConfig::Memory(_) => None,
CacheBackendConfig::Redis(_) if service.kind() != "redis" => {
Some("facade and native backend types must match")
}
CacheBackendConfig::Redis(config) if service.topology() != Some(&config.topology) => {
Some("facade and native backend topologies must match")
}
CacheBackendConfig::Redis(config) => (service.namespace()
!= config.namespace.as_deref())
.then_some("facade and native backend namespaces must match"),
}
}
}
#[inline(never)]
fn project_memory(backend: &Bound<'_, PyAny>) -> PyResult<MemoryCacheConfig> {
let max_size_kib = backend.getattr("max_size_per_item")?.extract::<usize>()?;
Ok(MemoryCacheConfig {
default_ttl: duration(backend.getattr("default_ttl")?.extract::<f64>()?)?,
capacity: backend.getattr("max_size_in_memory")?.extract::<usize>()?,
max_entry_bytes: max_size_kib
.checked_mul(1024)
.ok_or_else(|| PyValueError::new_err("memory cache item limit is too large"))?,
})
}
#[inline(never)]
fn project_redis(
backend: &Bound<'_, PyAny>,
) -> PyResult<Result<RedisCacheConfig, UnsupportedCacheConfig>> {
let source = backend.getattr("redis_kwargs")?.cast_into::<PyDict>()?;
if has_value(&source, "sentinel_nodes")? {
return Ok(Err(UnsupportedCacheConfig::RedisTopology));
}
for key in ["credential_provider", "redis_connect_func"] {
if has_value(&source, key)? {
return Ok(Err(UnsupportedCacheConfig::RedisCredentials));
}
}
if has_value(&source, "connection_pool")? {
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
}
for key in [
"retry",
"retry_on_error",
"socket_keepalive_options",
"unix_socket_path",
"cache",
"cache_config",
"event_dispatcher",
"ssl_ca_path",
"ssl_password",
"ssl_min_version",
"ssl_ciphers",
"ssl_validate_ocsp",
"ssl_validate_ocsp_stapled",
"ssl_ocsp_context",
"ssl_ocsp_expected_cert",
] {
if has_value(&source, key)? {
return Ok(Err(UnsupportedCacheConfig::RedisOption));
}
}
for key in ["retry_on_timeout", "single_connection_client"] {
if optional_coerced_bool(&source, key)?.unwrap_or(false) {
return Ok(Err(UnsupportedCacheConfig::RedisOption));
}
}
let client = backend.getattr("redis_client")?;
let projection = if has_value(&source, "startup_nodes")? {
project_cluster_client(&source, &client)?
} else {
project_standalone_client(&client)?
};
let RedisClientProjection {
topology,
host,
port,
pool_size,
resolved,
tls,
} = match projection {
Ok(projection) => projection,
Err(reason) => return Ok(Err(reason)),
};
if has_value(&resolved, "credential_provider")? {
return Ok(Err(UnsupportedCacheConfig::RedisCredentials));
}
let protocol = match optional_i64(&resolved, "protocol")?.unwrap_or(2) {
2 => RedisProtocol::Resp2,
3 => RedisProtocol::Resp3,
_ => return Err(PyValueError::new_err("unsupported Redis protocol version")),
};
let health_check_interval =
duration(optional_f64(&resolved, "health_check_interval")?.unwrap_or(0.0))?;
Ok(Ok(RedisCacheConfig {
default_ttl: duration(backend.getattr("default_ttl")?.extract::<f64>()?)?,
namespace: optional_attribute_string(backend, "namespace")?,
flush_size: backend.getattr("redis_flush_size")?.extract::<usize>()?,
topology,
connection: RedisConnectionConfig {
host,
port,
database: optional_i64(&resolved, "db")?.unwrap_or(0),
username: optional_dict_string(&resolved, "username")?,
password: optional_dict_string(&resolved, "password")?,
protocol,
pool_size,
read_timeout: optional_dict_duration(&resolved, "socket_timeout")?,
connect_timeout: optional_dict_duration(&resolved, "socket_connect_timeout")?,
socket_keepalive: optional_bool(&resolved, "socket_keepalive")?,
health_check_interval,
client_name: optional_dict_string(&resolved, "client_name")?,
tls,
},
}))
}
#[inline(never)]
fn project_standalone_client<'py>(
client: &Bound<'py, PyAny>,
) -> PyResult<Result<RedisClientProjection<'py>, UnsupportedCacheConfig>> {
let pool = client.getattr("connection_pool")?;
if !instance_class_is(&pool, "redis.connection", "ConnectionPool")? {
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
}
let resolved = pool.getattr("connection_kwargs")?.cast_into::<PyDict>()?;
if has_value(&resolved, "redis_connect_func")? {
return Ok(Err(UnsupportedCacheConfig::RedisCredentials));
}
let connection_class = resolved
.get_item("connection_class")?
.unwrap_or(pool.getattr("connection_class")?);
let tls = if class_is(&connection_class, "redis.connection", "Connection")? {
None
} else if class_is(&connection_class, "redis.connection", "SSLConnection")? {
Some(project_tls(&resolved)?)
} else {
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
};
Ok(Ok(RedisClientProjection {
topology: RedisTopology::Standalone,
host: required_string(&resolved, "host")?,
port: port(required_i64(&resolved, "port")?)?,
pool_size: pool.getattr("max_connections")?.extract::<usize>()?,
resolved,
tls,
}))
}
#[inline(never)]
fn project_cluster_client<'py>(
source: &Bound<'py, PyDict>,
client: &Bound<'py, PyAny>,
) -> PyResult<Result<RedisClientProjection<'py>, UnsupportedCacheConfig>> {
let Some(startup_nodes) = startup_nodes(source)? else {
return Ok(Err(UnsupportedCacheConfig::RedisTopology));
};
if !instance_class_is(client, "redis.cluster", "RedisCluster")? {
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
}
let nodes = client.getattr("nodes_manager")?;
if !class_is(
&nodes.getattr("connection_pool_class")?,
"redis.connection",
"ConnectionPool",
)? {
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
}
let resolved = nodes.getattr("connection_kwargs")?.cast_into::<PyDict>()?;
if let Some(connect) = resolved.get_item("redis_connect_func")?
&& !connect.is_none()
{
let own_hook = connect
.getattr("__self__")
.is_ok_and(|owner| owner.is(client))
&& connect
.getattr("__func__")
.and_then(|function| Ok(function.is(&client.get_type().getattr("on_connect")?)))
.unwrap_or(false);
if !own_hook {
return Ok(Err(UnsupportedCacheConfig::RedisCredentials));
}
}
let tls = if optional_bool(&resolved, "ssl")?.unwrap_or(false) {
Some(project_tls(&resolved)?)
} else {
None
};
let first = &startup_nodes[0];
Ok(Ok(RedisClientProjection {
host: first.host.clone(),
port: first.port,
pool_size: optional_i64(&resolved, "max_connections")?
.map(|value| {
usize::try_from(value).map_err(|_| PyValueError::new_err("invalid Redis pool size"))
})
.transpose()?
.unwrap_or(REDIS_PY_DEFAULT_MAX_CONNECTIONS),
topology: RedisTopology::Cluster { startup_nodes },
resolved,
tls,
}))
}
#[inline(never)]
fn startup_nodes(source: &Bound<'_, PyDict>) -> PyResult<Option<Vec<RedisNode>>> {
let Some(nodes) = source.get_item("startup_nodes")? else {
return Ok(None);
};
let Ok(nodes) = nodes.cast_into::<PyList>() else {
return Ok(None);
};
if nodes.is_empty() {
return Ok(None);
}
let mut parsed = Vec::with_capacity(nodes.len());
for node in nodes.iter() {
let Ok(node) = node.cast_into::<PyDict>() else {
return Ok(None);
};
if node.len() != 2 || !has_value(&node, "host")? || !has_value(&node, "port")? {
return Ok(None);
}
let (Ok(host), Ok(port)) = (
required_string(&node, "host"),
required_i64(&node, "port").and_then(port),
) else {
return Ok(None);
};
parsed.push(RedisNode { host, port });
}
Ok(Some(parsed))
}
#[inline(never)]
fn port(value: i64) -> PyResult<u16> {
u16::try_from(value).map_err(|_| PyValueError::new_err("invalid Redis port"))
}
#[inline(never)]
fn project_tls(values: &Bound<'_, PyDict>) -> PyResult<RedisTlsConfig> {
Ok(RedisTlsConfig {
certificate_requirement: certificate_requirement(values)?,
check_hostname: optional_bool(values, "ssl_check_hostname")?.unwrap_or(false),
ca_certificate: optional_dict_string(values, "ssl_ca_certs")?,
ca_data: optional_dict_string(values, "ssl_ca_data")?,
client_certificate: optional_dict_string(values, "ssl_certfile")?,
client_key: optional_dict_string(values, "ssl_keyfile")?,
})
}
#[inline(never)]
fn certificate_requirement(values: &Bound<'_, PyDict>) -> PyResult<CertificateRequirement> {
let Some(value) = values.get_item("ssl_cert_reqs")? else {
return Ok(CertificateRequirement::Required);
};
if value.is_none() {
return Ok(CertificateRequirement::Required);
}
if let Ok(number) = value.extract::<i32>() {
return match number {
0 => Ok(CertificateRequirement::None),
1 => Ok(CertificateRequirement::Optional),
2 => Ok(CertificateRequirement::Required),
_ => Err(PyValueError::new_err(
"invalid Redis TLS certificate requirement",
)),
};
}
let text = value.str()?;
let text = text.to_str()?;
if text.eq_ignore_ascii_case("none") || text.eq_ignore_ascii_case("cert_none") {
return Ok(CertificateRequirement::None);
}
if text.eq_ignore_ascii_case("optional") || text.eq_ignore_ascii_case("cert_optional") {
return Ok(CertificateRequirement::Optional);
}
if text.eq_ignore_ascii_case("required") || text.eq_ignore_ascii_case("cert_required") {
return Ok(CertificateRequirement::Required);
}
Err(PyValueError::new_err(
"invalid Redis TLS certificate requirement",
))
}
#[inline(never)]
fn instance_class_is(value: &Bound<'_, PyAny>, module: &str, name: &str) -> PyResult<bool> {
class_is(value.get_type().as_any(), module, name)
}
#[inline(never)]
fn class_is(value: &Bound<'_, PyAny>, module: &str, name: &str) -> PyResult<bool> {
Ok(value
.getattr("__module__")?
.cast_into::<PyString>()?
.to_str()?
== module
&& value
.getattr("__qualname__")?
.cast_into::<PyString>()?
.to_str()?
== name)
}
#[inline(never)]
fn optional_duration(value: Bound<'_, PyAny>) -> PyResult<Option<Duration>> {
value.extract::<Option<f64>>()?.map(duration).transpose()
}
#[inline(never)]
fn optional_attribute_string(value: &Bound<'_, PyAny>, name: &str) -> PyResult<Option<String>> {
match value.getattr(name) {
Ok(value) => optional_string(value),
Err(error) if error.is_instance_of::<pyo3::exceptions::PyAttributeError>(value.py()) => {
Ok(None)
}
Err(error) => Err(error),
}
}
#[inline(never)]
fn optional_string(value: Bound<'_, PyAny>) -> PyResult<Option<String>> {
Ok(value
.extract::<Option<String>>()?
.filter(|value| !value.is_empty()))
}
#[inline(never)]
fn has_value(values: &Bound<'_, PyDict>, key: &str) -> PyResult<bool> {
Ok(values.get_item(key)?.is_some_and(|value| !value.is_none()))
}
#[inline(never)]
fn required_string(values: &Bound<'_, PyDict>, key: &str) -> PyResult<String> {
values
.get_item(key)?
.ok_or_else(|| PyTypeError::new_err("Redis connection is incomplete"))?
.extract::<String>()
}
#[inline(never)]
fn required_i64(values: &Bound<'_, PyDict>, key: &str) -> PyResult<i64> {
values
.get_item(key)?
.ok_or_else(|| PyTypeError::new_err("Redis connection is incomplete"))?
.extract::<i64>()
}
#[inline(never)]
fn optional_dict_string(values: &Bound<'_, PyDict>, key: &str) -> PyResult<Option<String>> {
match values.get_item(key)? {
Some(value) if !value.is_none() => optional_string(value),
_ => Ok(None),
}
}
#[inline(never)]
fn optional_f64(values: &Bound<'_, PyDict>, key: &str) -> PyResult<Option<f64>> {
match values.get_item(key)? {
Some(value) => value.extract::<Option<f64>>(),
None => Ok(None),
}
}
#[inline(never)]
fn optional_i64(values: &Bound<'_, PyDict>, key: &str) -> PyResult<Option<i64>> {
match values.get_item(key)? {
Some(value) => value.extract::<Option<i64>>(),
None => Ok(None),
}
}
#[inline(never)]
fn optional_bool(values: &Bound<'_, PyDict>, key: &str) -> PyResult<Option<bool>> {
match values.get_item(key)? {
Some(value) => value.extract::<Option<bool>>(),
None => Ok(None),
}
}
#[inline(never)]
fn optional_coerced_bool(values: &Bound<'_, PyDict>, key: &str) -> PyResult<Option<bool>> {
let Some(value) = values.get_item(key)? else {
return Ok(None);
};
if value.is_none() {
return Ok(None);
}
if let Ok(text) = value.extract::<String>() {
return Ok(Some(
text == "1" || text.eq_ignore_ascii_case("true") || text.eq_ignore_ascii_case("yes"),
));
}
value.extract::<bool>().map(Some)
}
#[inline(never)]
fn optional_dict_duration(values: &Bound<'_, PyDict>, key: &str) -> PyResult<Option<Duration>> {
optional_f64(values, key)?.map(duration).transpose()
}
#[cfg(test)]
mod tests {
use std::ffi::CString;
use pyo3::{prelude::*, types::PyDict};
use litellm_cache_redis::{RedisNode, RedisTopology};
use super::{
CacheBackendConfig, CacheConfigProjection, CertificateRequirement, NativeCacheConfig,
RedisProtocol,
};
use crate::cache::native::NativeResponseCache;
fn cluster_facade<'py>(py: Python<'py>, startup_nodes: &str, hook: &str) -> Bound<'py, PyAny> {
facade(
py,
&format!(
"RedisCluster = type('RedisCluster', (), {{'__module__': 'redis.cluster', 'on_connect': lambda self, connection: None}})\n\
client = RedisCluster()\n\
client.nodes_manager = SimpleNamespace(connection_pool_class=ConnectionPool, connection_kwargs={{'password': 'secret', 'redis_connect_func': {hook}, 'protocol': 3, 'ssl': True, 'ssl_cert_reqs': 'none'}})\n\
backend = SimpleNamespace(default_ttl=120, namespace='team', redis_flush_size=100, redis_kwargs={{'startup_nodes': {startup_nodes}, 'password': 'secret'}}, redis_client=client)\n\
facade = SimpleNamespace(type='redis', mode='default-on', ttl=None, namespace='team', supported_call_types=None, redis_flush_size=100, semantic_cache_scope='key', cache=backend)"
),
)
}
fn facade<'py>(py: Python<'py>, body: &str) -> Bound<'py, PyAny> {
let locals = PyDict::new(py);
py.run(
&CString::new(format!(
"from types import SimpleNamespace\n\
ConnectionPool = type('ConnectionPool', (), {{'__module__': 'redis.connection'}})\n\
Connection = type('Connection', (), {{'__module__': 'redis.connection'}})\n\
SSLConnection = type('SSLConnection', (), {{'__module__': 'redis.connection'}})\n\
{body}"
))
.unwrap(),
None,
Some(&locals),
)
.unwrap();
locals.get_item("facade").unwrap().unwrap()
}
#[test]
fn projects_effective_memory_configuration() {
Python::initialize();
Python::attach(|py| {
let facade = facade(
py,
"backend = SimpleNamespace(default_ttl=913, max_size_in_memory=37, max_size_per_item=8)\n\
facade = SimpleNamespace(type='local', mode='default-on', ttl=11.5, namespace=None, supported_call_types=['completion'], redis_flush_size=None, semantic_cache_scope='key', cache=backend)",
);
let CacheConfigProjection::Native(config) =
NativeCacheConfig::project(&facade).unwrap()
else {
panic!("memory cache should be supported");
};
assert_eq!(
config.policy.ttl.unwrap(),
std::time::Duration::from_secs_f64(11.5)
);
let CacheBackendConfig::Memory(memory) = config.backend else {
panic!("expected memory configuration");
};
assert_eq!(memory.default_ttl, std::time::Duration::from_secs(913));
assert_eq!(memory.capacity, 37);
assert_eq!(memory.max_entry_bytes, 8192);
let matching =
NativeResponseCache::memory(37, std::time::Duration::from_secs(913), 8192);
let mismatched =
NativeResponseCache::memory(37, std::time::Duration::from_secs(913), 8191);
let matching_config = NativeCacheConfig {
policy: config.policy,
backend: CacheBackendConfig::Memory(memory),
};
assert_eq!(matching_config.service_mismatch(&matching), None);
assert_eq!(
matching_config.service_mismatch(&mismatched),
Some("facade and native backend item limits must match")
);
});
}
#[test]
fn projects_resolved_redis_tls_configuration() {
Python::initialize();
Python::attach(|py| {
let facade = facade(
py,
"pool = ConnectionPool()\n\
pool.connection_class = SSLConnection\n\
pool.max_connections = 29\n\
pool.connection_kwargs = {'host': 'cache.internal', 'port': 6380, 'db': 4, 'username': 'user', 'password': 'secret', 'protocol': 3, 'socket_timeout': 7.5, 'socket_connect_timeout': 2, 'socket_keepalive': True, 'health_check_interval': 15, 'client_name': 'litellm', 'ssl_cert_reqs': 'optional', 'ssl_check_hostname': True, 'ssl_ca_certs': '/ca.pem', 'ssl_ca_data': 'CA DATA', 'ssl_certfile': '/client.pem', 'ssl_keyfile': '/client.key'}\n\
client = SimpleNamespace(connection_pool=pool)\n\
backend = SimpleNamespace(default_ttl=777, namespace='team', redis_flush_size=31, redis_kwargs={}, redis_client=client)\n\
facade = SimpleNamespace(type='redis', mode='default-off', ttl=None, namespace='team', supported_call_types=None, redis_flush_size=31, semantic_cache_scope='key', cache=backend)",
);
let CacheConfigProjection::Native(config) =
NativeCacheConfig::project(&facade).unwrap()
else {
panic!("Redis cache should be supported");
};
let CacheBackendConfig::Redis(redis) = config.backend else {
panic!("expected Redis configuration");
};
assert_eq!(redis.default_ttl, std::time::Duration::from_secs(777));
assert_eq!(redis.namespace.as_deref(), Some("team"));
assert_eq!(redis.flush_size, 31);
assert_eq!(redis.connection.host, "cache.internal");
assert_eq!(redis.connection.port, 6380);
assert_eq!(redis.connection.database, 4);
assert_eq!(redis.connection.protocol, RedisProtocol::Resp3);
assert_eq!(redis.connection.pool_size, 29);
let tls = redis.connection.tls.unwrap();
assert_eq!(
tls.certificate_requirement,
CertificateRequirement::Optional
);
assert!(tls.check_hostname);
assert_eq!(tls.ca_certificate.as_deref(), Some("/ca.pem"));
assert_eq!(tls.ca_data.as_deref(), Some("CA DATA"));
assert_eq!(tls.client_certificate.as_deref(), Some("/client.pem"));
assert_eq!(tls.client_key.as_deref(), Some("/client.key"));
});
}
#[test]
fn dynamic_redis_auth_stays_on_python() {
Python::initialize();
Python::attach(|py| {
let facade = facade(
py,
"backend = SimpleNamespace(redis_kwargs={'credential_provider': object()})\n\
facade = SimpleNamespace(type='redis', mode='default-on', ttl=None, namespace=None, supported_call_types=[], redis_flush_size=None, semantic_cache_scope='key', cache=backend)",
);
let CacheConfigProjection::Unsupported(reason) =
NativeCacheConfig::project(&facade).unwrap()
else {
panic!("dynamic authentication must stay on Python");
};
assert_eq!(reason.message(), "native Redis credentials require Python");
});
}
#[test]
fn projects_cluster_startup_nodes_as_redis_topology() {
Python::initialize();
Python::attach(|py| {
let facade = cluster_facade(
py,
"[{'host': 'node-a', 'port': 7000}, {'host': 'node-b', 'port': 7001}]",
"client.on_connect",
);
let CacheConfigProjection::Native(config) =
NativeCacheConfig::project(&facade).unwrap()
else {
panic!("cluster startup nodes should project natively");
};
let CacheBackendConfig::Redis(redis) = &config.backend else {
panic!("expected Redis configuration");
};
let expected = RedisTopology::Cluster {
startup_nodes: vec![
RedisNode {
host: "node-a".into(),
port: 7000,
},
RedisNode {
host: "node-b".into(),
port: 7001,
},
],
};
assert_eq!(redis.topology, expected);
assert_eq!(redis.connection.host, "node-a");
assert_eq!(redis.connection.port, 7000);
assert_eq!(redis.connection.password.as_deref(), Some("secret"));
assert_eq!(redis.connection.protocol, RedisProtocol::Resp3);
assert_eq!(
redis
.connection
.tls
.as_ref()
.unwrap()
.certificate_requirement,
CertificateRequirement::None
);
});
}
#[test]
fn malformed_startup_nodes_and_foreign_connect_hooks_stay_on_python() {
Python::initialize();
Python::attach(|py| {
for (startup_nodes, hook, message) in [
(
"[{'host': 'node-a', 'port': 7000, 'server_type': 'primary'}]",
"client.on_connect",
"native Redis topology is not implemented",
),
(
"[{'host': 'node-a', 'port': 'seven'}]",
"client.on_connect",
"native Redis topology is not implemented",
),
(
"[]",
"client.on_connect",
"native Redis topology is not implemented",
),
(
"[{'host': 'node-a', 'port': 7000}]",
"lambda connection: None",
"native Redis credentials require Python",
),
] {
let facade = cluster_facade(py, startup_nodes, hook);
let CacheConfigProjection::Unsupported(reason) =
NativeCacheConfig::project(&facade).unwrap()
else {
panic!("{startup_nodes} with {hook} must stay on Python");
};
assert_eq!(reason.message(), message, "{startup_nodes} with {hook}");
}
});
}
}

View file

@ -0,0 +1,330 @@
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::{
config::{CacheConfigProjection, NativeCacheConfig},
handle::CacheTestHandle,
native::NativeResponseCache,
};
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>,
attributes: RedisPoolAttributes,
}
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,
};
pub(super) struct FacadeGuard {
outer: ObjectGuard,
backend: ObjectGuard,
redis_pool: Option<RedisPoolGuard>,
}
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)? || !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>, attributes: RedisPoolAttributes) -> PyResult<Self> {
let pool = backend.getattr("redis_client")?.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: Self::max_connections(&pool, &attributes)?,
attributes,
})
}
fn max_connections(
pool: &Bound<'_, PyAny>,
attributes: &RedisPoolAttributes,
) -> PyResult<Option<usize>> {
attributes
.max_connections
.map(|name| pool.getattr(name)?.extract::<usize>())
.transpose()
}
fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult<bool> {
let pool = backend
.getattr("redis_client")?
.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::max_connections(&pool, &self.attributes)?
&& 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 FacadeGuard {
pub(super) fn capture(
py: Python<'_>,
facade: &Bound<'_, PyAny>,
service: &NativeResponseCache,
) -> PyResult<Self> {
let kind = service.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!(service.topology(), Some(RedisTopology::Cluster { .. }));
let (module, name, cache_kind) = match (kind, cluster) {
("memory", _) => ("litellm.caching.in_memory_cache", "InMemoryCache", "local"),
("redis", false) => ("litellm.caching.redis_cache", "RedisCache", "redis"),
("redis", true) => (
"litellm.caching.redis_cluster_cache",
"RedisClusterCache",
"redis",
),
_ => unreachable!(),
};
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));
}
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",
],
)?,
redis_pool: match (kind, cluster) {
("redis", false) => Some(RedisPoolGuard::capture(&backend, STANDALONE_POOL)?),
("redis", true) => Some(RedisPoolGuard::capture(&backend, CLUSTER_POOL)?),
_ => None,
},
})
}
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);
}
match &self.redis_pool {
Some(guard) => guard.matches(py, &backend),
None => Ok(true),
}
}
pub(super) fn traverse(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
self.outer.traverse(&visit)?;
self.backend.traverse(&visit)?;
if let Some(guard) = &self.redis_pool {
guard.traverse(&visit)?;
}
Ok(())
}
}
pub(super) fn resolve(
py: Python<'_>,
facade: &Bound<'_, PyAny>,
) -> PyResult<Option<NativeResponseCache>> {
let Ok(dict) = facade
.getattr("__dict__")
.and_then(|dict| dict.cast_into::<PyDict>().map_err(Into::into))
else {
return Ok(None);
};
let Some(handle) = dict.get_item("_native_cache_handle")? else {
return Ok(None);
};
let Ok(handle) = handle.extract::<PyRef<'_, CacheTestHandle>>() else {
return Ok(None);
};
let Some(guard) = &handle.guard else {
return Ok(None);
};
if !guard.matches(py, facade).unwrap_or(false) {
return Ok(None);
}
handle.service().map(Some)
}

View file

@ -0,0 +1,18 @@
use litellm_host_python::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>> {
let future = py
.import("asyncio")?
.call_method0("get_running_loop")?
.call_method0("create_future")?;
future.call_method1("set_result", (to_py(py, value)?,))?;
Ok(future)
}

View file

@ -0,0 +1,97 @@
use litellm_cache_redis::{RedisNode, RedisTopology};
use litellm_host_python::release_gil;
use pyo3::{PyTraverseError, PyVisit, exceptions::PyRuntimeError, prelude::*};
use super::{cache_error, facade::FacadeGuard, native::NativeResponseCache, request::duration};
#[pyclass(frozen, name = "_CacheTestHandle")]
pub(crate) struct CacheTestHandle {
service: NativeResponseCache,
pub(super) guard: Option<FacadeGuard>,
pid: u32,
}
impl CacheTestHandle {
pub(super) fn service(&self) -> PyResult<NativeResponseCache> {
if self.pid != std::process::id() {
return Err(PyRuntimeError::new_err(
"native cache handles must be recreated after fork",
));
}
Ok(self.service.clone())
}
}
#[pymethods]
impl CacheTestHandle {
#[staticmethod]
#[pyo3(signature = (*, capacity=200, ttl_seconds=600.0, max_entry_bytes=1048576))]
fn memory(capacity: usize, ttl_seconds: f64, max_entry_bytes: usize) -> PyResult<Self> {
Ok(Self {
service: NativeResponseCache::memory(capacity, duration(ttl_seconds)?, max_entry_bytes),
guard: None,
pid: std::process::id(),
})
}
#[staticmethod]
#[pyo3(signature = (url, *, ttl_seconds=60.0, namespace=None, startup_nodes=None))]
fn redis(
py: Python<'_>,
url: String,
ttl_seconds: f64,
namespace: Option<String>,
startup_nodes: Option<Vec<(String, u16)>>,
) -> PyResult<Self> {
let ttl = Some(duration(ttl_seconds)?);
let topology = match startup_nodes {
None => RedisTopology::Standalone,
Some(nodes) => RedisTopology::Cluster {
startup_nodes: nodes
.into_iter()
.map(|(host, port)| RedisNode { host, port })
.collect(),
},
};
let service = release_gil(py, move || {
NativeResponseCache::redis(&url, &topology, ttl, namespace)
})
.map_err(cache_error)?;
Ok(Self {
service,
guard: None,
pid: std::process::id(),
})
}
#[getter]
fn backend(&self) -> &'static str {
self.service.kind()
}
fn _bind_facade(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<()> {
let service = self.service()?;
let guard = FacadeGuard::capture(py, facade, &service)?;
let service = service.with_redis_flush_size(
facade
.getattr("redis_flush_size")?
.extract::<Option<usize>>()?,
);
let handle = Py::new(
py,
Self {
service,
guard: Some(guard),
pid: self.pid,
},
)?;
facade.setattr("_native_cache_handle", handle)
}
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
if let Some(guard) = &self.guard {
guard.traverse(visit)?;
}
Ok(())
}
}

View file

@ -0,0 +1,26 @@
mod binding;
mod callback;
mod config;
mod facade;
mod future;
mod handle;
mod native;
mod request;
mod resolver;
use litellm_cache::Error;
use pyo3::{
exceptions::{PyRuntimeError, PyValueError},
prelude::*,
};
pub(crate) use self::{
binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheTestResolver,
};
fn cache_error(error: Error) -> PyErr {
match error {
Error::InvalidEntry => PyValueError::new_err(error.to_string()),
_ => PyRuntimeError::new_err(error.to_string()),
}
}

View file

@ -0,0 +1,207 @@
use std::{sync::Arc, time::Duration};
use litellm_cache::{CacheCodec, CacheConnectionResult, Error};
use litellm_cache_memory::InMemoryCache;
use litellm_cache_redis::{RedisCache, RedisTopology};
use litellm_cache_response::{
CacheEntry, PartialHits, ResponseCache, ResponseCacheCodec, ResponseCacheRequest, WriteBuffer,
};
use serde_json::Value;
#[derive(Clone)]
pub(super) enum NativeResponseCache {
Memory(Arc<ResponseCache<InMemoryCache<CacheEntry>>>),
Redis {
cache: Arc<ResponseCache<RedisCache<ResponseCacheCodec>>>,
buffer: Option<Arc<WriteBuffer>>,
},
}
impl NativeResponseCache {
pub fn memory(capacity: usize, ttl: Duration, max_entry_bytes: usize) -> Self {
Self::Memory(Arc::new(ResponseCache::new(Arc::new(
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())
})),
super::request::now,
),
))))
}
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);
Ok(Self::Redis {
cache: Arc::new(ResponseCache::new(Arc::new(backend))),
buffer: None,
})
}
}
impl NativeResponseCache {
pub fn kind(&self) -> &'static str {
match self {
Self::Memory(_) => "memory",
Self::Redis { .. } => "redis",
}
}
pub fn default_ttl(&self) -> Option<Duration> {
match self {
Self::Memory(cache) => cache.default_ttl(),
Self::Redis { cache, .. } => cache.default_ttl(),
}
}
pub fn namespace(&self) -> Option<&str> {
match self {
Self::Memory(_) => None,
Self::Redis { cache, .. } => cache.backend().namespace(),
}
}
pub fn topology(&self) -> Option<&RedisTopology> {
match self {
Self::Memory(_) => None,
Self::Redis { cache, .. } => Some(cache.backend().topology()),
}
}
pub fn capacity(&self) -> Option<usize> {
match self {
Self::Memory(cache) => Some(cache.backend().max_size_in_memory()),
Self::Redis { .. } => None,
}
}
pub fn max_entry_bytes(&self) -> Option<usize> {
match self {
Self::Memory(cache) => cache.backend().max_entry_bytes(),
Self::Redis { .. } => None,
}
}
pub fn with_redis_flush_size(self, flush_size: Option<usize>) -> Self {
match self {
Self::Redis { cache, .. } => Self::Redis {
cache,
buffer: flush_size.map(|flush_size| Arc::new(WriteBuffer::new(flush_size))),
},
memory => memory,
}
}
pub fn lookup(
&self,
request: &ResponseCacheRequest,
now: Duration,
) -> Result<Option<Value>, Error> {
match self {
Self::Memory(cache) => cache.lookup(request, now),
Self::Redis { cache, .. } => cache.lookup(request, now),
}
}
pub fn store(
&self,
request: &ResponseCacheRequest,
response: Value,
now: Duration,
) -> Result<(), Error> {
match self {
Self::Memory(cache) => cache.store(request, response, now),
Self::Redis { cache, .. } => cache.store(request, response, now),
}
}
pub fn lookup_batch(
&self,
requests: &[ResponseCacheRequest],
now: Duration,
) -> Result<PartialHits, Error> {
match self {
Self::Memory(cache) => cache.lookup_batch(requests, now),
Self::Redis { cache, .. } => cache.lookup_batch(requests, now),
}
}
pub async fn async_lookup(
&self,
request: &ResponseCacheRequest,
now: Duration,
) -> Result<Option<Value>, Error> {
match self {
Self::Memory(cache) => cache.async_lookup(request, now).await,
Self::Redis { cache, .. } => cache.async_lookup(request, now).await,
}
}
pub async fn async_store(
&self,
request: &ResponseCacheRequest,
response: Value,
now: Duration,
) -> Result<(), Error> {
match self {
Self::Memory(cache) => cache.async_store(request, response, now).await,
Self::Redis {
cache,
buffer: None,
} => cache.async_store(request, response, now).await,
Self::Redis {
cache,
buffer: Some(buffer),
} => buffer.async_store(cache, request, response, now).await,
}
}
pub async fn async_lookup_batch(
&self,
requests: &[ResponseCacheRequest],
now: Duration,
) -> Result<PartialHits, Error> {
match self {
Self::Memory(cache) => cache.async_lookup_batch(requests, now).await,
Self::Redis { cache, .. } => cache.async_lookup_batch(requests, now).await,
}
}
pub async fn async_store_batch(
&self,
entries: Vec<(ResponseCacheRequest, Value)>,
now: Duration,
) -> Result<(), Error> {
match self {
Self::Memory(cache) => cache.async_store_batch(entries, now).await,
Self::Redis { cache, .. } => cache.async_store_batch(entries, now).await,
}
}
pub async fn async_flush(&self) -> Result<(), Error> {
match self {
Self::Memory(cache) => cache.async_flush().await,
Self::Redis { cache, buffer } => {
if let Some(buffer) = buffer {
buffer.clear()?;
}
cache.async_flush().await
}
}
}
pub async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
match self {
Self::Memory(cache) => cache.test_connection().await,
Self::Redis { cache, .. } => cache.test_connection().await,
}
}
}

View file

@ -0,0 +1,48 @@
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use litellm_cache_response::{CacheControls, CacheKeyInput, ResponseCacheRequest};
use litellm_host_python::from_py;
use pyo3::{exceptions::PyValueError, prelude::*};
use serde::Deserialize;
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct RequestInput {
key: CacheKeyInput,
controls: Option<CacheControls>,
ttl_seconds: Option<f64>,
max_age_seconds: Option<f64>,
}
pub(super) fn request(value: &Bound<'_, PyAny>) -> PyResult<ResponseCacheRequest> {
let input: RequestInput = from_py(value)?;
request_input(input)
}
fn request_input(input: RequestInput) -> PyResult<ResponseCacheRequest> {
let mut request = ResponseCacheRequest::new(input.key);
if let Some(controls) = input.controls {
request.controls = controls;
}
request.context.ttl = input.ttl_seconds.map(duration).transpose()?;
request.max_age = input.max_age_seconds.map(duration).transpose()?;
Ok(request)
}
pub(super) fn requests(value: &Bound<'_, PyAny>) -> PyResult<Vec<ResponseCacheRequest>> {
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(super) fn now() -> Duration {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
}

View file

@ -0,0 +1,39 @@
use pyo3::{PyTraverseError, PyVisit, prelude::*};
use super::{
binding::{CacheBinding, ResolvedCache},
callback::PythonCallback,
facade,
handle::CacheTestHandle,
};
#[pyclass(frozen, name = "_CacheTestResolver")]
pub(crate) struct CacheTestResolver {
namespace: Py<PyAny>,
}
#[pymethods]
impl CacheTestResolver {
#[new]
fn new(namespace: Py<PyAny>) -> Self {
Self { namespace }
}
pub(crate) fn resolve(&self, py: Python<'_>) -> PyResult<ResolvedCache> {
let object = self.namespace.bind(py).getattr("cache")?;
let binding = if object.is_none() {
CacheBinding::Disabled
} else if let Ok(handle) = object.extract::<PyRef<'_, CacheTestHandle>>() {
CacheBinding::Native(handle.service()?)
} else if let Some(service) = facade::resolve(py, &object)? {
CacheBinding::Native(service)
} else {
CacheBinding::PythonCallback(PythonCallback::new(object.unbind()))
};
Ok(ResolvedCache::new(binding))
}
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.namespace)
}
}

View file

@ -0,0 +1,231 @@
use std::collections::BTreeSet;
use litellm_core_utils::serde_compat::parse_str_bool;
use litellm_http::SslVerify;
use pyo3::{
exceptions::{PyAttributeError, PyRuntimeError, PyValueError},
prelude::*,
types::{PyBool, PyString},
};
#[derive(Debug)]
pub(crate) enum ProjectionError {
Python(PyErr),
InvalidConfiguration(String),
UnsupportedLiveObject(String),
InternalSchemaFailure(String),
}
impl From<PyErr> for ProjectionError {
fn from(error: PyErr) -> Self {
Self::Python(error)
}
}
impl From<ProjectionError> for PyErr {
fn from(error: ProjectionError) -> Self {
match error {
ProjectionError::Python(error) => error,
ProjectionError::InvalidConfiguration(message)
| ProjectionError::UnsupportedLiveObject(message) => PyValueError::new_err(message),
ProjectionError::InternalSchemaFailure(message) => PyRuntimeError::new_err(message),
}
}
}
pub(crate) struct Truthy(pub bool);
pub(crate) struct ExactTrue(pub bool);
pub(crate) struct StrBool(pub Option<bool>);
pub(crate) struct OptionalStrictString(pub Option<String>);
pub(crate) struct FalsyOptionalString(pub Option<String>);
pub(crate) struct TuningString(pub Option<String>);
pub(crate) struct StringCollection(pub Vec<String>);
pub(crate) struct SslVerifyInput(pub Option<SslVerify>);
pub(crate) struct Field<'py> {
path: &'static str,
value: Bound<'py, PyAny>,
}
impl<'py> Field<'py> {
pub(crate) fn new(path: &'static str, value: Bound<'py, PyAny>) -> Self {
Self { path, value }
}
pub(crate) fn read(
snapshot: &Bound<'py, PyAny>,
path: &'static str,
) -> Result<Self, ProjectionError> {
let name = path.rsplit('.').next().unwrap_or(path);
match snapshot.getattr(name) {
Ok(value) => Ok(Self::new(path, value)),
Err(error) if error.is_instance_of::<PyAttributeError>(snapshot.py()) => {
match Self::missing_field(snapshot, name) {
Ok(true) => Err(ProjectionError::InternalSchemaFailure(format!(
"{path}: missing snapshot field"
))),
_ => Err(error.into()),
}
}
Err(error) => Err(error.into()),
}
}
fn missing_field(snapshot: &Bound<'_, PyAny>, name: &str) -> PyResult<bool> {
let py = snapshot.py();
let object = py.import("builtins")?.getattr("object")?;
let missing = object.call0()?;
let lookup = py.import("inspect")?.getattr("getattr_static")?;
let declared = lookup.call1((snapshot, name, &missing))?;
let fallback = lookup.call1((snapshot.get_type(), "__getattr__", &missing))?;
let getter = lookup.call1((snapshot.get_type(), "__getattribute__"))?;
Ok(declared.is(&missing)
&& fallback.is(&missing)
&& getter.is(object.getattr("__getattribute__")?))
}
fn expected(&self, expected: &'static str) -> Result<String, ProjectionError> {
Ok(format!(
"{}: expected {expected}, got {}",
self.path,
self.value.get_type().name()?
))
}
fn invalid(&self, expected: &'static str) -> ProjectionError {
match self.expected(expected) {
Ok(message) => ProjectionError::InvalidConfiguration(message),
Err(error) => error,
}
}
pub(crate) fn truthy(&self) -> Result<Truthy, ProjectionError> {
Ok(Truthy(self.value.is_truthy()?))
}
pub(crate) fn exact_true(&self) -> ExactTrue {
ExactTrue(self.value.is(PyBool::new(self.value.py(), true)))
}
pub(crate) fn strict_string(&self) -> Result<String, ProjectionError> {
let value = self
.value
.cast::<PyString>()
.map_err(|_| self.invalid("a string"))?;
Ok(value.to_str()?.to_owned())
}
pub(crate) fn schema_string(&self) -> Result<String, ProjectionError> {
if !self.value.is_instance_of::<PyString>() {
return Err(ProjectionError::InternalSchemaFailure(
self.expected("a string")?,
));
}
self.strict_string()
}
pub(crate) fn schema_bool(&self) -> Result<bool, ProjectionError> {
if !self.value.is_instance_of::<PyBool>() {
return Err(ProjectionError::InternalSchemaFailure(
self.expected("a Boolean")?,
));
}
Ok(self.exact_true().0)
}
pub(crate) fn str_bool(&self) -> Result<StrBool, ProjectionError> {
if self.value.is_none() {
return Ok(StrBool(None));
}
Ok(StrBool(parse_str_bool(&self.strict_string()?)))
}
pub(crate) fn optional_strict_string(&self) -> Result<OptionalStrictString, ProjectionError> {
if self.value.is_none() {
return Ok(OptionalStrictString(None));
}
self.strict_string().map(Some).map(OptionalStrictString)
}
pub(crate) fn falsy_optional_string(&self) -> Result<FalsyOptionalString, ProjectionError> {
if !self.truthy()?.0 {
return Ok(FalsyOptionalString(None));
}
self.strict_string().map(Some).map(FalsyOptionalString)
}
pub(crate) fn tuning_string(&self) -> Result<TuningString, ProjectionError> {
if !self.truthy()?.0 || !self.value.is_instance_of::<PyString>() {
return Ok(TuningString(None));
}
self.strict_string().map(Some).map(TuningString)
}
pub(crate) fn string_collection(&self) -> Result<StringCollection, ProjectionError> {
if !self.truthy()?.0 {
return Ok(StringCollection(Vec::new()));
}
if self.value.is_instance_of::<PyString>() {
return self
.strict_string()
.map(|value| StringCollection(vec![value]));
}
let values = self
.value
.try_iter()?
.filter_map(|item| {
let member = match item {
Ok(value) => Self::new(self.path, value),
Err(error) => return Some(Err(error.into())),
};
match member.truthy() {
Ok(Truthy(false)) => None,
Ok(Truthy(true)) => Some(member.strict_string()),
Err(error) => Some(Err(error)),
}
})
.collect::<Result<Vec<_>, ProjectionError>>()?;
Ok(StringCollection(values))
}
pub(crate) fn host_collection(&self) -> Result<StringCollection, ProjectionError> {
let values = self
.string_collection()?
.0
.into_iter()
.map(|host| litellm_http::media::normalize_host(&host))
.collect::<BTreeSet<_>>();
Ok(StringCollection(values.into_iter().collect()))
}
pub(crate) fn ssl_verify(&self) -> Result<SslVerifyInput, ProjectionError> {
if self.value.is_none() {
return Ok(SslVerifyInput(None));
}
if self.value.is_instance_of::<PyBool>() {
return Ok(SslVerifyInput(Some(if self.exact_true().0 {
SslVerify::Enabled
} else {
SslVerify::Disabled
})));
}
if self.value.is_instance_of::<PyString>() {
let parsed = match self.str_bool()?.0 {
Some(true) => SslVerify::Enabled,
Some(false) => SslVerify::Disabled,
None => SslVerify::CaBundle(self.strict_string()?.into()),
};
return Ok(SslVerifyInput(Some(parsed)));
}
let context = self.value.py().import("ssl")?.getattr("SSLContext")?;
if self.value.is_instance(&context)? {
return Err(ProjectionError::UnsupportedLiveObject(self.expected(
"a Boolean, Boolean string, CA path, or None; live SSLContext is unsupported",
)?));
}
Err(self.invalid("a Boolean, Boolean string, CA path, or None"))
}
}
#[cfg(test)]
mod tests;

View file

@ -0,0 +1,372 @@
use std::ffi::CString;
use pyo3::{
exceptions::{PyLookupError, PyRuntimeError, PyValueError},
types::PyDict,
};
use rstest::rstest;
use super::*;
fn evaluate<'py>(py: Python<'py>, source: &str) -> Bound<'py, PyAny> {
py.eval(&CString::new(source).unwrap(), None, None).unwrap()
}
#[rstest]
#[case("None", false, false)]
#[case("False", false, false)]
#[case("True", true, true)]
#[case("0", false, false)]
#[case("1", true, false)]
#[case("''", false, false)]
#[case("'false'", true, false)]
#[case("[]", false, false)]
#[case("[0]", true, false)]
#[case("{}", false, false)]
#[case("object()", true, false)]
fn boolean_operations_have_distinct_python_semantics(
#[case] source: &str,
#[case] truth: bool,
#[case] exact: bool,
) {
Python::initialize();
Python::attach(|py| {
let value = evaluate(py, source);
let field = Field::new("test.flag", value.clone());
assert_eq!(field.truthy().unwrap().0, truth);
assert_eq!(field.exact_true().0, exact);
assert_eq!(
field.truthy().unwrap().0,
py.import("builtins")
.unwrap()
.getattr("bool")
.unwrap()
.call1((value,))
.unwrap()
.extract::<bool>()
.unwrap()
);
});
}
#[rstest]
#[case("None", Ok(None), Ok(None), Ok(None))]
#[case("''", Ok(Some("")), Ok(None), Ok(None))]
#[case(
"' value '",
Ok(Some(" value ")),
Ok(Some(" value ")),
Ok(Some(" value "))
)]
#[case("[]", Err(()), Ok(None), Ok(None))]
#[case("0", Err(()), Ok(None), Ok(None))]
#[case("1", Err(()), Err(()), Ok(None))]
#[case("object()", Err(()), Err(()), Ok(None))]
fn string_operations_do_not_conflate_absence_and_type_checks(
#[case] source: &str,
#[case] strict: Result<Option<&str>, ()>,
#[case] fallback: Result<Option<&str>, ()>,
#[case] tuning: Result<Option<&str>, ()>,
) {
Python::initialize();
Python::attach(|py| {
let field = Field::new("test.string", evaluate(py, source));
let owned =
|expected: Result<Option<&str>, ()>| expected.map(|value| value.map(str::to_owned));
assert_eq!(
field
.optional_strict_string()
.map(|value| value.0)
.map_err(|_| ()),
owned(strict)
);
assert_eq!(
field
.falsy_optional_string()
.map(|value| value.0)
.map_err(|_| ()),
owned(fallback)
);
assert_eq!(
field.tuning_string().map(|value| value.0).map_err(|_| ()),
owned(tuning)
);
});
}
#[rstest]
#[case("None", None)]
#[case("' True '", Some(true))]
#[case("' fAlSe '", Some(false))]
#[case("'yes'", None)]
#[case("'1'", None)]
#[case("'unknown'", None)]
fn string_boolean_tokens_remain_separate_from_truthiness(
#[case] source: &str,
#[case] expected: Option<bool>,
) {
Python::initialize();
Python::attach(|py| {
assert_eq!(
Field::new("test.flag", evaluate(py, source))
.str_bool()
.unwrap()
.0,
expected
);
});
}
#[rstest]
#[case("'EXAMPLE.TEST.'", vec!["example.test"])]
#[case("['B.test', '', None, 0, [], 'A.test.', 'b.test']", vec!["a.test", "b.test"])]
#[case("('B.test', 'a.test')", vec!["a.test", "b.test"])]
#[case("{'B.test', 'a.test'}", vec!["a.test", "b.test"])]
#[case("(host for host in ['B.test', 'a.test'])", vec!["a.test", "b.test"])]
#[case("None", vec![])]
#[case("False", vec![])]
fn host_collection_is_owned_normalized_and_deterministic(
#[case] source: &str,
#[case] expected: Vec<&str>,
) {
Python::initialize();
Python::attach(|py| {
assert_eq!(
Field::new("url_policy.user_url_allowed_hosts", evaluate(py, source))
.host_collection()
.unwrap()
.0,
expected
);
});
}
#[test]
fn protocol_errors_preserve_exception_identity_traceback_cause_and_context() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"
failure = LookupError('protocol failed')
cause = ValueError('cause')
context = RuntimeError('context')
def fail():
try:
raise context
except RuntimeError:
raise failure from cause
class Bool:
def __bool__(self): return fail()
class Length:
def __len__(self): return fail()
class Iter:
def __iter__(self): return fail()
class Next:
def __iter__(self): return self
def __next__(self): return fail()
class Descriptor:
@property
def flag(self): return fail()
values = (Bool(), Length(), Iter(), Next(), [Bool()])
descriptor = Descriptor()
",
Some(&locals),
Some(&locals),
)
.unwrap();
let values = locals.get_item("values").unwrap().unwrap();
for value in values.try_iter().unwrap() {
let error = Field::new("test.flag", value.unwrap())
.host_collection()
.err()
.unwrap();
let error = PyErr::from(error);
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
assert!(error.is_instance_of::<PyLookupError>(py));
assert!(error.traceback(py).is_some());
assert!(
error
.value(py)
.getattr("__cause__")
.unwrap()
.is(locals.get_item("cause").unwrap().unwrap())
);
assert!(
error
.value(py)
.getattr("__context__")
.unwrap()
.is(locals.get_item("context").unwrap().unwrap())
);
}
let error = Field::read(
&locals.get_item("descriptor").unwrap().unwrap(),
"test.flag",
)
.err()
.unwrap();
assert!(
PyErr::from(error)
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
});
}
#[test]
fn identity_and_string_contents_do_not_invoke_unrelated_protocols() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"
class Hostile:
def __bool__(self): raise AssertionError('bool called')
def __eq__(self, other): raise AssertionError('eq called')
def __str__(self): raise AssertionError('str called')
class Text(str):
def __str__(self): raise AssertionError('str called')
def strip(self): raise AssertionError('strip called')
def lower(self): raise AssertionError('lower called')
hostile = Hostile()
text = Text(' False ')
",
Some(&locals),
Some(&locals),
)
.unwrap();
let hostile = Field::new("test.flag", locals.get_item("hostile").unwrap().unwrap());
assert!(!hostile.exact_true().0);
assert!(matches!(
hostile.strict_string(),
Err(ProjectionError::InvalidConfiguration(_))
));
let text = Field::new("test.flag", locals.get_item("text").unwrap().unwrap());
assert_eq!(text.strict_string().unwrap(), " False ");
assert_eq!(text.str_bool().unwrap().0, Some(false));
});
}
#[test]
fn missing_snapshot_fields_and_descriptor_attribute_errors_are_distinct() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"
failure = AttributeError('descriptor failed')
class Snapshot:
@property
def flag(self): raise failure
snapshot = Snapshot()
class Dynamic:
def __getattr__(self, name): raise failure
class Intercepted:
def __getattribute__(self, name): raise failure
dynamic = Dynamic()
intercepted = Intercepted()
",
Some(&locals),
Some(&locals),
)
.unwrap();
let snapshot = locals.get_item("snapshot").unwrap().unwrap();
let descriptor = PyErr::from(Field::read(&snapshot, "test.flag").err().unwrap());
assert!(
descriptor
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
for name in ["dynamic", "intercepted"] {
let value = locals.get_item(name).unwrap().unwrap();
let error = PyErr::from(Field::read(&value, "test.flag").err().unwrap());
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
}
let missing = PyErr::from(Field::read(&snapshot, "test.missing").err().unwrap());
assert!(missing.is_instance_of::<PyRuntimeError>(py));
assert!(missing.to_string().contains("test.missing"));
});
}
#[test]
fn configuration_errors_name_fields_without_exposing_values() {
Python::initialize();
Python::attach(|py| {
for source in [
"{'secret': 'do-not-print'}",
"['host.test', {'secret': 'do-not-print'}]",
] {
let field = Field::new("test.setting", evaluate(py, source));
let error = PyErr::from(field.falsy_optional_string().err().unwrap());
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("test.setting"));
assert!(!error.to_string().contains("do-not-print"));
}
let hosts = Field::new(
"url_policy.user_url_allowed_hosts",
evaluate(py, "['host.test', 1]"),
);
assert!(matches!(
hosts.host_collection(),
Err(ProjectionError::InvalidConfiguration(_))
));
assert!(matches!(
Field::new("test.flag", evaluate(py, "1")).str_bool(),
Err(ProjectionError::InvalidConfiguration(_))
));
});
}
#[test]
fn projection_releases_the_source_collection() {
Python::initialize();
Python::attach(|py| {
let source = evaluate(py, "['A.test']");
let projected = Field::new("test.hosts", source.clone())
.host_collection()
.unwrap()
.0;
source.call_method1("append", ("b.test",)).unwrap();
assert_eq!(projected, ["a.test"]);
assert_eq!(
Field::new("test.hosts", source)
.host_collection()
.unwrap()
.0,
["a.test", "b.test"]
);
});
}
#[rstest]
#[case("True", Some(true))]
#[case("False", Some(false))]
#[case("1", None)]
#[case("None", None)]
#[case("[]", None)]
fn accessor_booleans_are_strict_schema_values(
#[case] source: &str,
#[case] expected: Option<bool>,
) {
Python::initialize();
Python::attach(|py| {
let result = Field::new("secret_manager.readable", evaluate(py, source)).schema_bool();
match expected {
Some(expected) => assert_eq!(result.unwrap(), expected),
None => {
let error = PyErr::from(result.unwrap_err());
assert!(error.is_instance_of::<PyRuntimeError>(py));
assert!(error.to_string().contains("secret_manager.readable"));
}
}
});
}

View file

@ -7,12 +7,12 @@ use std::{
use litellm_core_utils::settings::ProcessEnvironment;
use litellm_http::{
HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer, Resolution, SslVerify,
Unsupported,
TlsSource, Unsupported,
media::{PublicDnsResolver, UrlPolicy},
};
use pyo3::{prelude::*, types::PyDict};
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
use crate::{errors::RustBridgeDeclined, python_settings::PythonSettings};
use crate::{coercion::Field, python_settings::PythonSettings};
static POOL: LazyLock<HttpClientPool> =
LazyLock::new(|| HttpClientPool::new(Arc::new(PublicDnsResolver)));
@ -41,6 +41,30 @@ pub(crate) fn call_config(
Ok(resolution.config)
}
pub(crate) fn client_error(error: litellm_http::Error) -> PyErr {
match error {
litellm_http::Error::Read {
tls_source: TlsSource::ClientIdentity,
..
}
| litellm_http::Error::InvalidPem {
tls_source: TlsSource::ClientIdentity,
..
} => PyValueError::new_err(
"http_settings.ssl_certificate: expected a readable PEM certificate and private key",
),
litellm_http::Error::Read {
tls_source: TlsSource::CaBundle,
..
}
| litellm_http::Error::InvalidPem {
tls_source: TlsSource::CaBundle,
..
} => PyValueError::new_err("http_settings.ssl_verify: expected a readable PEM CA bundle"),
_ => PyValueError::new_err("http_settings: native HTTP client configuration is invalid"),
}
}
fn unreported(
reported: &Mutex<HashSet<Unsupported>>,
unsupported: Vec<Unsupported>,
@ -53,25 +77,25 @@ fn unreported(
}
pub(crate) fn url_policy(py: Python<'_>) -> PyResult<UrlPolicy> {
let policy: PythonUrlPolicy =
PythonSettings::UrlPolicy
.read(py)?
.extract()
.map_err(|error: PyErr| {
RustBridgeDeclined::new_err(format!(
"litellm URL policy cannot be used by the Rust route: {error}"
))
})?;
project_url_policy(&PythonSettings::UrlPolicy.read(py)?)
}
fn project_url_policy(value: &Bound<'_, PyAny>) -> PyResult<UrlPolicy> {
Ok(UrlPolicy {
validate: policy.user_url_validation,
allowed_hosts: policy.user_url_allowed_hosts,
validate: Field::read(value, "url_policy.user_url_validation")?
.truthy()?
.0,
allowed_hosts: Field::read(value, "url_policy.user_url_allowed_hosts")?
.host_collection()?
.0,
})
}
fn call_ssl_verify(kwargs: &Bound<'_, PyDict>) -> PyResult<Option<SslVerify>> {
Ok(kwargs
.get_item("ssl_verify")?
.and_then(|value| ssl_verify(&value)))
match kwargs.get_item("ssl_verify")? {
Some(value) => Ok(Field::new("request.ssl_verify", value).ssl_verify()?.0),
None => Ok(None),
}
}
fn for_call(call_ssl_verify: Option<SslVerify>, asynchronous: bool) -> HttpSettingsLayer {
@ -82,64 +106,47 @@ fn for_call(call_ssl_verify: Option<SslVerify>, asynchronous: bool) -> HttpSetti
}
}
#[derive(FromPyObject)]
struct PythonUrlPolicy {
user_url_validation: bool,
user_url_allowed_hosts: Vec<String>,
}
#[derive(FromPyObject)]
struct PythonHttpSettings<'py> {
ssl_verify: Bound<'py, PyAny>,
ssl_certificate: Option<String>,
ssl_security_level: Option<String>,
ssl_ecdh_curve: Option<String>,
force_ipv4: bool,
http2: bool,
aiohttp_trust_env: bool,
disable_aiohttp_trust_env: bool,
disable_aiohttp_transport: bool,
user_agent: String,
}
fn configured(value: &Bound<'_, PyAny>) -> PyResult<HttpSettingsLayer> {
let python: PythonHttpSettings = value.extract().map_err(|error: PyErr| {
RustBridgeDeclined::new_err(format!(
"litellm HTTP settings cannot be used by the Rust route: {error}"
))
})?;
Ok(HttpSettingsLayer {
ssl_verify: ssl_verify(&python.ssl_verify),
ssl_certificate: python.ssl_certificate.map(PathBuf::from),
ssl_security_level: python.ssl_security_level,
ssl_ecdh_curve: python.ssl_ecdh_curve,
force_ipv4: Some(python.force_ipv4),
http2: Some(python.http2),
aiohttp_trust_env: Some(python.aiohttp_trust_env),
disable_aiohttp_trust_env: Some(python.disable_aiohttp_trust_env),
disable_aiohttp_transport: Some(python.disable_aiohttp_transport),
user_agent: Some(python.user_agent),
ssl_verify: Field::read(value, "http_settings.ssl_verify")?
.ssl_verify()?
.0,
ssl_certificate: Field::read(value, "http_settings.ssl_certificate")?
.optional_strict_string()?
.0
.map(PathBuf::from),
ssl_security_level: Field::read(value, "http_settings.ssl_security_level")?
.tuning_string()?
.0,
ssl_ecdh_curve: Field::read(value, "http_settings.ssl_ecdh_curve")?
.tuning_string()?
.0,
force_ipv4: Some(Field::read(value, "http_settings.force_ipv4")?.truthy()?.0),
http2: Some(Field::read(value, "http_settings.http2")?.exact_true().0),
aiohttp_trust_env: Some(
Field::read(value, "http_settings.aiohttp_trust_env")?
.truthy()?
.0,
),
disable_aiohttp_trust_env: Some(
Field::read(value, "http_settings.disable_aiohttp_trust_env")?
.truthy()?
.0,
),
disable_aiohttp_transport: Some(
Field::read(value, "http_settings.disable_aiohttp_transport")?
.exact_true()
.0,
),
user_agent: Some(Field::read(value, "http_settings.user_agent")?.schema_string()?),
..HttpSettingsLayer::default()
})
}
fn ssl_verify(value: &Bound<'_, PyAny>) -> Option<SslVerify> {
if let Ok(enabled) = value.extract::<bool>() {
return Some(if enabled {
SslVerify::Enabled
} else {
SslVerify::Disabled
});
}
value
.extract::<String>()
.ok()
.map(|path| SslVerify::parse(&path))
}
#[cfg(test)]
mod tests {
use litellm_http::Verify;
use pyo3::exceptions::PyRuntimeError;
use rstest::rstest;
use super::*;
@ -163,7 +170,7 @@ defaults = dict(
user_agent='litellm/test',
)
defaults.update(dict({overrides}))
settings = types.SimpleNamespace(**{{name: defaults[name] for name in json.loads(contract)['http_settings']}})
settings = types.SimpleNamespace(**{{name: defaults[name] for name in json.loads(contract)['http_settings']['fields']}})
"
);
let locals = PyDict::new(py);
@ -189,6 +196,33 @@ settings = types.SimpleNamespace(**{{name: defaults[name] for name in json.loads
});
}
#[test]
fn client_error_uses_tls_source_when_paths_match() {
Python::initialize();
Python::attach(|py| {
let path = PathBuf::from("/shared.pem");
let ca_error = client_error(litellm_http::Error::InvalidPem {
path: path.clone(),
message: "invalid".into(),
tls_source: TlsSource::CaBundle,
});
assert_eq!(
ca_error.to_string(),
"ValueError: http_settings.ssl_verify: expected a readable PEM CA bundle"
);
let client_error = client_error(litellm_http::Error::InvalidPem {
path,
message: "invalid".into(),
tls_source: TlsSource::ClientIdentity,
});
assert!(client_error.is_instance_of::<PyValueError>(py));
assert_eq!(
client_error.to_string(),
"ValueError: http_settings.ssl_certificate: expected a readable PEM certificate and private key"
);
});
}
#[test]
fn python_settings_flow_into_the_configured_layer() {
Python::initialize();
@ -259,12 +293,16 @@ user_agent='litellm/9.9.9',
});
}
#[test]
fn ssl_context_global_is_ignored_so_environment_and_defaults_apply() {
#[rstest]
#[case("ssl_verify=object()")]
#[case("ssl_verify=__import__('ssl').SSLContext(__import__('ssl').PROTOCOL_TLS_CLIENT)")]
#[case("ssl_certificate=1")]
fn invalid_http_configuration_is_terminal(#[case] overrides: &str) {
Python::initialize();
Python::attach(|py| {
let layer = configured(&python_settings(py, "ssl_verify=object()")).unwrap();
assert_eq!(layer.ssl_verify, None);
let error = configured(&python_settings(py, overrides)).unwrap_err();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("http_settings.ssl_"));
});
}
@ -281,11 +319,21 @@ user_agent='litellm/9.9.9',
}
#[test]
fn mistyped_python_settings_decline_instead_of_raising() {
fn mutable_globals_use_their_consumer_operations() {
Python::initialize();
Python::attach(|py| {
let error = configured(&python_settings(py, "force_ipv4='yes'")).unwrap_err();
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
let layer = configured(&python_settings(py,
"force_ipv4='yes', http2=1, disable_aiohttp_transport=1, aiohttp_trust_env=[1], disable_aiohttp_trust_env=[], ssl_security_level=1, ssl_ecdh_curve=[]"
)).unwrap();
assert_eq!(layer.force_ipv4, Some(true));
assert_eq!(layer.http2, Some(false));
assert_eq!(layer.disable_aiohttp_transport, Some(false));
assert_eq!(layer.aiohttp_trust_env, Some(true));
assert_eq!(layer.disable_aiohttp_trust_env, Some(false));
assert_eq!(layer.ssl_security_level, None);
assert_eq!(layer.ssl_ecdh_curve, None);
let error = configured(&python_settings(py, "user_agent=1")).unwrap_err();
assert!(error.is_instance_of::<PyRuntimeError>(py));
});
}
@ -323,17 +371,36 @@ user_agent='litellm/9.9.9',
}
#[test]
fn live_ssl_context_argument_is_ignored_so_the_configured_value_applies() {
fn live_ssl_context_argument_raises_instead_of_using_another_layer() {
Python::initialize();
Python::attach(|py| {
let kwargs = PyDict::new(py);
kwargs
.set_item("ssl_verify", py.eval(c"object()", None, None).unwrap())
let ssl = py.import("ssl").unwrap();
let context = ssl
.getattr("SSLContext")
.unwrap()
.call1((ssl.getattr("PROTOCOL_TLS_CLIENT").unwrap(),))
.unwrap();
let call = for_call(call_ssl_verify(&kwargs).unwrap(), true);
let settings =
HttpSettings::from_layers([call, configured_ssl_verify(SslVerify::Disabled)]);
assert_eq!(settings.ssl_verify, Some(SslVerify::Disabled));
kwargs.set_item("ssl_verify", context).unwrap();
let error = call_ssl_verify(&kwargs).unwrap_err();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("request.ssl_verify"));
assert!(error.to_string().contains("SSLContext"));
});
}
#[test]
fn url_policy_uses_truthiness_and_normalized_owned_hosts() {
Python::initialize();
Python::attach(|py| {
let value = py.eval(c"__import__('types').SimpleNamespace(user_url_validation=[], user_url_allowed_hosts=['B.test', 'a.test.', 'b.test'])", None, None).unwrap();
assert_eq!(
project_url_policy(&value).unwrap(),
UrlPolicy {
validate: false,
allowed_hosts: vec!["a.test".into(), "b.test".into()],
}
);
});
}

View file

@ -1,3 +1,5 @@
mod cache;
mod coercion;
mod credentials;
mod diagnostics;
mod errors;
@ -9,6 +11,7 @@ mod token_counter;
#[pymodule(gil_used = true)]
mod _native {
use crate::cache::{CacheTestHandle, CacheTestResolver, ResolvedCache};
#[cfg(feature = "panic-test")]
#[pymodule_export]
use crate::diagnostics::_panic_for_test;
@ -32,6 +35,16 @@ mod _native {
use crate::token_counter::TokenCounter;
#[pymodule_export]
use litellm_host_python::{ForkedAfterNativeRuntimeStarted, ProcessReservedForForking};
use pyo3::{prelude::*, types::PyModule};
#[pymodule_init]
fn init(module: &Bound<'_, PyModule>) -> PyResult<()> {
let py = module.py();
let dict = module.dict();
dict.set_item("_CacheTestHandle", py.get_type::<CacheTestHandle>())?;
dict.set_item("_CacheTestResolver", py.get_type::<CacheTestResolver>())?;
dict.set_item("_CacheTestBinding", py.get_type::<ResolvedCache>())
}
}
use pyo3::prelude::*;

View file

@ -172,6 +172,60 @@ mod tests {
request_input_sources(&kwargs, names.iter().copied())
}
#[serde_with::serde_as]
#[derive(Debug, serde::Deserialize, serde::Serialize, PartialEq)]
struct Numbers {
#[serde_as(deserialize_as = "Option<Vec<litellm_core_utils::serde_compat::LaxI64>>")]
integers: Option<Vec<i64>>,
#[serde_as(deserialize_as = "Option<litellm_core_utils::serde_compat::FiniteF64>")]
float: Option<f64>,
}
#[test]
fn numeric_adapters_agree_across_json_and_python_boundaries() {
Python::initialize();
Python::attach(|py| {
for input in [
json!({}),
json!({"integers": null, "float": null}),
json!({"integers": [i64::MIN, i64::MAX, "9007199254740993.0", " +1_000.00 ", true, 3.0], "float": " 1.25 "}),
json!({"integers": [u64::MAX]}),
json!({"integers": ["1.0000000000000001"]}),
json!({"integers": [2.5]}),
json!({"float": "NaN"}),
json!({"float": "inf"}),
json!({"float": "1e999"}),
json!({"float": true}),
json!({"float": u64::MAX}),
] {
let expected = serde_json::from_value::<Numbers>(input.clone());
let python = litellm_host_python::to_py(py, &input).unwrap();
let actual = from_py::<Numbers>(python.bind(py));
match (expected, actual) {
(Ok(expected), Ok(actual)) => {
assert_eq!(actual, expected);
let serialized = litellm_host_python::to_py(py, &actual).unwrap();
assert_eq!(
from_py::<Value>(serialized.bind(py)).unwrap(),
serde_json::to_value(expected).unwrap()
);
}
(Err(_), Err(_)) => {}
mismatch => panic!("boundary mismatch for {input}: {mismatch:?}"),
}
}
for source in [
c"{'float': float('nan')}",
c"{'float': float('inf')}",
c"{'integers': [float('inf')]}",
c"{'integers': [2 ** 100]}",
] {
let value = py.eval(source, None, None).unwrap();
assert!(from_py::<Numbers>(&value).is_err());
}
});
}
#[test]
fn argument_converters_keep_nested_values_and_accept_explicit_none() {
Python::initialize();

View file

@ -43,32 +43,204 @@ pub(crate) const CONTRACT: &str = include_str!("../python_settings.json");
#[cfg(test)]
mod tests {
use std::{collections::BTreeSet, ffi::CString};
use pyo3::{prelude::*, types::PyDict};
use super::{CONTRACT, PythonSettings};
use pyo3::prelude::*;
use serde_json::{Value, json};
struct SettingSpec {
group: &'static str,
name: &'static str,
adapter: &'static str,
precedence: &'static str,
sensitive: bool,
shapes: &'static [&'static str],
unsupported_live: Option<&'static str>,
}
const SETTINGS: &[SettingSpec] = &[
SettingSpec {
group: "http_settings",
name: "ssl_verify",
adapter: "SslVerifyInput",
precedence: "module_global",
sensitive: false,
shapes: &["none", "bool", "str"],
unsupported_live: Some("configuration_error"),
},
SettingSpec {
group: "http_settings",
name: "ssl_certificate",
adapter: "OptionalStrictString",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "ssl_security_level",
adapter: "TuningString",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "ssl_ecdh_curve",
adapter: "TuningString",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "force_ipv4",
adapter: "Truthy",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "http2",
adapter: "ExactTrue",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "aiohttp_trust_env",
adapter: "Truthy",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "disable_aiohttp_trust_env",
adapter: "Truthy",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "disable_aiohttp_transport",
adapter: "ExactTrue",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "http_settings",
name: "user_agent",
adapter: "StrictString",
precedence: "accessor",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "url_policy",
name: "user_url_validation",
adapter: "Truthy",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "url_policy",
name: "user_url_allowed_hosts",
adapter: "HostCollection",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "provider_defaults",
name: "vertex_project",
adapter: "FalsyOptionalString",
precedence: "module_global",
sensitive: true,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "provider_defaults",
name: "vertex_location",
adapter: "FalsyOptionalString",
precedence: "module_global",
sensitive: true,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "provider_defaults",
name: "enable_azure_ad_token_refresh",
adapter: "ExactTrue",
precedence: "module_global",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
SettingSpec {
group: "secret_manager",
name: "readable",
adapter: "StrictBool",
precedence: "accessor",
sensitive: false,
shapes: &[],
unsupported_live: None,
},
];
#[test]
fn every_settings_group_is_in_the_python_contract() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
locals.set_item("contract", CONTRACT).unwrap();
let source = CString::new("import json\nkeys = list(json.loads(contract))").unwrap();
py.run(&source, Some(&locals), Some(&locals)).unwrap();
let declared: BTreeSet<String> = locals
.get_item("keys")
fn settings_manifest_matches_the_semantic_contract() {
pyo3::Python::initialize();
let manifest: Value = pyo3::Python::attach(|py| {
let value = py
.import("json")
.unwrap()
.unwrap()
.extract::<Vec<String>>()
.unwrap()
.into_iter()
.collect();
let read: BTreeSet<String> = PythonSettings::ALL
.map(|group| group.name().to_owned())
.into();
assert_eq!(read, declared);
.call_method1("loads", (CONTRACT,))
.unwrap();
litellm_host_python::from_py(&value).unwrap()
});
let expected: serde_json::Map<String, Value> = PythonSettings::ALL
.into_iter()
.map(|group| {
let fields: serde_json::Map<String, Value> = SETTINGS
.iter()
.filter(|spec| spec.group == group.name())
.map(|spec| {
(
spec.name.to_owned(),
json!({
"adapter": spec.adapter,
"required": true,
"precedence": spec.precedence,
"sensitive": spec.sensitive,
"shapes": spec.shapes,
"unsupported_live": spec.unsupported_live,
}),
)
})
.collect();
(
group.name().to_owned(),
json!({"version": 1, "fields": fields}),
)
})
.collect();
assert_eq!(manifest, Value::Object(expected));
}
}

View file

@ -19,7 +19,7 @@ use pyo3::{
types::{PyDict, PyTuple},
};
use crate::{errors::RustBridgeDeclined, http, python_settings::PythonSettings};
use crate::{coercion::Field, errors::RustBridgeDeclined, http, python_settings::PythonSettings};
const SURFACE: LegacySurface = LegacySurface {
call_type: "ocr",
@ -51,7 +51,7 @@ fn run_ocr(
ocr_settings(py)?,
secrets,
)
.map_err(|error| RustBridgeDeclined::new_err(error.to_string()))?;
.map_err(http::client_error)?;
run_legacy_call(
py,
if asynchronous { ASYNC_SURFACE } else { SURFACE },
@ -62,14 +62,8 @@ fn run_ocr(
)
}
#[derive(FromPyObject)]
struct PythonSecretManager {
readable: bool,
}
fn process_environment_secrets(secret_manager: &Bound<'_, PyAny>) -> PyResult<Secrets> {
let manager: PythonSecretManager = secret_manager.extract()?;
if manager.readable {
if Field::read(secret_manager, "secret_manager.readable")?.schema_bool()? {
return Err(RustBridgeDeclined::new_err(
"a readable secret manager is configured and the Rust route only reads the process environment",
));
@ -77,26 +71,24 @@ fn process_environment_secrets(secret_manager: &Bound<'_, PyAny>) -> PyResult<Se
Ok(Arc::new(ProcessEnvironment))
}
#[derive(FromPyObject)]
struct PythonProviderDefaults {
vertex_project: Option<String>,
vertex_location: Option<String>,
enable_azure_ad_token_refresh: Option<bool>,
fn ocr_settings(py: Python<'_>) -> PyResult<OcrSettings> {
project_provider_defaults(&PythonSettings::ProviderDefaults.read(py)?)
}
fn ocr_settings(py: Python<'_>) -> PyResult<OcrSettings> {
let defaults: PythonProviderDefaults = PythonSettings::ProviderDefaults
.read(py)?
.extract()
.map_err(|error: PyErr| {
RustBridgeDeclined::new_err(format!(
"litellm provider defaults cannot be used by the Rust route: {error}"
))
})?;
fn project_provider_defaults(value: &Bound<'_, PyAny>) -> PyResult<OcrSettings> {
Ok(OcrSettings {
vertex_project: defaults.vertex_project,
vertex_location: defaults.vertex_location,
enable_azure_ad_token_refresh: defaults.enable_azure_ad_token_refresh == Some(true),
vertex_project: Field::read(value, "provider_defaults.vertex_project")?
.falsy_optional_string()?
.0,
vertex_location: Field::read(value, "provider_defaults.vertex_location")?
.falsy_optional_string()?
.0,
enable_azure_ad_token_refresh: Field::read(
value,
"provider_defaults.enable_azure_ad_token_refresh",
)?
.exact_true()
.0,
..OcrSettings::from_environment(&ProcessEnvironment)
})
}
@ -140,6 +132,35 @@ mod tests {
locals.get_item("manager").unwrap().unwrap()
}
#[test]
fn provider_defaults_distinguish_falsey_values_and_exact_true() {
Python::initialize();
Python::attach(|py| {
let value = py.eval(c"__import__('types').SimpleNamespace(vertex_project=[], vertex_location=0, enable_azure_ad_token_refresh=1)", None, None).unwrap();
let projected = super::project_provider_defaults(&value).unwrap();
assert_eq!(projected.vertex_project, None);
assert_eq!(projected.vertex_location, None);
assert!(!projected.enable_azure_ad_token_refresh);
value.setattr("vertex_project", "project").unwrap();
value.setattr("vertex_location", "region").unwrap();
value
.setattr("enable_azure_ad_token_refresh", true)
.unwrap();
let next = super::project_provider_defaults(&value).unwrap();
assert_eq!(next.vertex_project.as_deref(), Some("project"));
assert_eq!(next.vertex_location.as_deref(), Some("region"));
assert!(next.enable_azure_ad_token_refresh);
value.setattr("vertex_project", 1).unwrap();
let error = super::project_provider_defaults(&value).err().unwrap();
assert!(error.is_instance_of::<pyo3::exceptions::PyValueError>(py));
assert!(
error
.to_string()
.contains("provider_defaults.vertex_project")
);
});
}
#[test]
fn a_readable_secret_manager_sends_the_call_back_to_python() {
Python::initialize();

View file

@ -0,0 +1,26 @@
[package]
name = "litellm-secrets-cyberark"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-secrets-types.workspace = true
litellm-core-utils.workspace = true
base64.workspace = true
moka.workspace = true
reqwest.workspace = true
serde_json.workspace = true
thiserror.workspace = true
veil.workspace = true
tracing = "0.1"
percent-encoding = "2.3"
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
rstest.workspace = true
tokio.workspace = true
wiremock = "0.6.5"
serde.workspace = true
serde_json.workspace = true

View file

@ -0,0 +1,27 @@
#[derive(thiserror::Error, veil::Redact)]
pub enum Error {
#[error("CyberArk Conjur HTTP request failed")]
Http(
#[from]
#[redact]
reqwest::Error,
),
#[error("CyberArk Conjur authentication returned HTTP {0}")]
AuthStatus(u16),
#[error("CyberArk Conjur returned HTTP {0}")]
Status(u16),
#[error(
"CyberArk credentials are missing: set CYBERARK_API_KEY or both CYBERARK_CLIENT_CERT and CYBERARK_CLIENT_KEY"
)]
MissingCredentials,
#[error("CyberArk client certificate could not be loaded")]
ClientCertificate,
#[error("invalid refresh interval")]
RefreshInterval,
#[error("invalid CyberArk Conjur endpoint")]
Endpoint,
#[error("CyberArk secret manager requires an enterprise license")]
EnterpriseRequired,
#[error(transparent)]
Operation(#[from] litellm_secrets_types::Error),
}

View file

@ -0,0 +1,7 @@
#![forbid(unsafe_code)]
mod error;
mod secret_manager;
pub use error::Error;
pub use secret_manager::{CyberArkSecretManager, DeleteOutcome};

View file

@ -0,0 +1,317 @@
use std::{fs, sync::Arc, time::Duration};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_core_utils::settings::Lookup;
use litellm_secrets_types::{BaseSecretManager, SecretValue, validate_secret_name};
use moka::future::Cache;
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, utf8_percent_encode};
use crate::Error;
const CYBERARK_API_BASE: &str = "CYBERARK_API_BASE";
const CYBERARK_ACCOUNT: &str = "CYBERARK_ACCOUNT";
const CYBERARK_USERNAME: &str = "CYBERARK_USERNAME";
const CYBERARK_API_KEY: &str = "CYBERARK_API_KEY";
const CYBERARK_CLIENT_CERT: &str = "CYBERARK_CLIENT_CERT";
const CYBERARK_CLIENT_KEY: &str = "CYBERARK_CLIENT_KEY";
const CYBERARK_SSL_VERIFY: &str = "CYBERARK_SSL_VERIFY";
const CYBERARK_REFRESH_INTERVAL: &str = "CYBERARK_REFRESH_INTERVAL";
const DEFAULT_API_BASE: &str = "http://127.0.0.1:8080";
const DEFAULT_ACCOUNT: &str = "default";
const DEFAULT_USERNAME: &str = "admin";
const DEFAULT_REFRESH_INTERVAL: Duration = Duration::from_secs(300);
const SECRET_NAME_SAFE: &AsciiSet = &NON_ALPHANUMERIC
.remove(b'-')
.remove(b'_')
.remove(b'.')
.remove(b'~');
#[derive(Clone)]
pub struct CyberArkSecretManager {
client: reqwest::Client,
endpoint: reqwest::Url,
account: String,
username: String,
api_key: SecretValue,
token: Cache<(), SecretValue>,
secrets: Cache<String, SecretValue>,
authentication_lock: Arc<tokio::sync::Mutex<()>>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DeleteOutcome {
NotSupported,
}
impl CyberArkSecretManager {
pub fn with_client(
client: reqwest::Client,
endpoint: reqwest::Url,
account: String,
username: String,
api_key: SecretValue,
refresh_interval: Option<Duration>,
) -> Self {
let endpoint = normalize_endpoint(endpoint);
let ttl = refresh_interval
.filter(|interval| !interval.is_zero())
.unwrap_or(DEFAULT_REFRESH_INTERVAL);
let token = Cache::builder().time_to_live(ttl).build();
let secrets = Cache::builder().time_to_live(ttl).build();
Self {
client,
endpoint,
account,
username,
api_key,
token,
secrets,
authentication_lock: Arc::new(tokio::sync::Mutex::new(())),
}
}
pub fn new(
environment: Arc<dyn Lookup + Send + Sync>,
enterprise_enabled: bool,
) -> Result<Self, Error> {
let api_key = environment.get(CYBERARK_API_KEY).unwrap_or_default();
let cert = environment.get(CYBERARK_CLIENT_CERT).unwrap_or_default();
let key = environment.get(CYBERARK_CLIENT_KEY).unwrap_or_default();
if api_key.is_empty() && (cert.is_empty() || key.is_empty()) {
return Err(Error::MissingCredentials);
}
if !enterprise_enabled {
return Err(Error::EnterpriseRequired);
}
let verify = environment
.get(CYBERARK_SSL_VERIFY)
.map(|value| !value.trim().eq_ignore_ascii_case("false"))
.unwrap_or(true);
let mut builder = reqwest::Client::builder();
if !verify {
tracing::warn!(
"CyberArk SSL verification is disabled. This is insecure and should only be used for testing with self-signed certificates."
);
builder = builder.danger_accept_invalid_certs(true);
}
if !cert.is_empty() && !key.is_empty() {
let certificate = fs::read(cert).map_err(|_| Error::ClientCertificate)?;
let private_key = fs::read(key).map_err(|_| Error::ClientCertificate)?;
let identity = reqwest::Identity::from_pem(&[certificate, private_key].concat())
.map_err(|_| Error::ClientCertificate)?;
builder = builder.identity(identity);
}
let client = builder.build()?;
let endpoint = reqwest::Url::parse(
&environment
.get(CYBERARK_API_BASE)
.unwrap_or_else(|| DEFAULT_API_BASE.to_owned()),
)
.map_err(|_| Error::Endpoint)?;
let account = environment
.get(CYBERARK_ACCOUNT)
.unwrap_or_else(|| DEFAULT_ACCOUNT.to_owned());
let username = environment
.get(CYBERARK_USERNAME)
.unwrap_or_else(|| DEFAULT_USERNAME.to_owned());
let refresh_interval = environment
.get(CYBERARK_REFRESH_INTERVAL)
.map(|value| {
value
.parse::<u64>()
.map(Duration::from_secs)
.map_err(|_| Error::RefreshInterval)
})
.transpose()?;
Ok(Self::with_client(
client,
endpoint,
account,
username,
SecretValue::new(api_key),
refresh_interval,
))
}
fn secret_url(&self, name: &str) -> Result<reqwest::Url, Error> {
let encoded = utf8_percent_encode(name, SECRET_NAME_SAFE);
self.endpoint
.join(&format!("secrets/{}/variable/{}", self.account, encoded))
.map_err(|_| Error::Endpoint)
}
async fn authenticate(&self) -> Result<SecretValue, Error> {
if let Some(token) = self.token.get(&()).await {
return Ok(token);
}
let _guard = self.authentication_lock.lock().await;
if let Some(token) = self.token.get(&()).await {
return Ok(token);
}
let url = self
.endpoint
.join(&format!(
"authn/{}/{}/authenticate",
self.account, self.username
))
.map_err(|_| Error::Endpoint)?;
let response = self
.client
.post(url)
.body(self.api_key.expose().to_owned())
.send()
.await?;
if !response.status().is_success() {
return Err(Error::AuthStatus(response.status().as_u16()));
}
let token = SecretValue::new(STANDARD.encode(response.text().await?));
self.token.insert((), token.clone()).await;
Ok(token)
}
async fn authorization_header(&self) -> Result<String, Error> {
Ok(format!(
"Token token=\"{}\"",
self.authenticate().await?.expose()
))
}
pub async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
if let Some(value) = self.secrets.get(name).await {
return Ok(Some(value));
}
let response = self
.client
.get(self.secret_url(name)?)
.header("Authorization", self.authorization_header().await?)
.send()
.await?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Ok(None);
}
if !response.status().is_success() {
return Err(Error::Status(response.status().as_u16()));
}
let value = SecretValue::new(response.text().await?);
self.secrets.insert(name.to_owned(), value.clone()).await;
Ok(Some(value))
}
pub async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
_description: Option<&str>,
) -> Result<(), Error> {
validate_secret_name(name)?;
self.ensure_variable_exists(name).await;
let response = self
.client
.post(self.secret_url(name)?)
.header("Authorization", self.authorization_header().await?)
.body(value.expose().to_owned())
.send()
.await?;
if !response.status().is_success() {
return Err(Error::Status(response.status().as_u16()));
}
self.secrets.insert(name.to_owned(), value.clone()).await;
Ok(())
}
async fn ensure_variable_exists(&self, name: &str) {
let policy_url = self
.endpoint
.join(&format!("policies/{}/policy/root", self.account));
let Ok(policy_url) = policy_url else {
tracing::warn!("Could not build CyberArk policy endpoint");
return;
};
let Ok(authorization) = self.authorization_header().await else {
tracing::warn!("Could not authenticate while ensuring CyberArk variable exists");
return;
};
let body = format!(
"- !variable {}\n",
serde_json::to_string(name).expect("serializing a string cannot fail")
);
let response = self
.client
.post(policy_url)
.header("Authorization", authorization)
.header("Content-Type", "application/x-yaml")
.body(body)
.send()
.await;
match response {
Ok(response) if response.status().is_success() => {}
Ok(response)
if matches!(
response.status(),
reqwest::StatusCode::CONFLICT | reqwest::StatusCode::UNPROCESSABLE_ENTITY
) =>
{
tracing::debug!(
"CyberArk variable policy already exists or conflicts: {}",
response.status()
);
}
Ok(response) => {
tracing::warn!(
"Could not ensure CyberArk variable exists: {}",
response.status()
);
}
Err(error) => {
tracing::warn!("Error ensuring CyberArk variable exists: {error}");
}
}
}
pub async fn async_delete_secret(
&self,
name: &str,
_recovery_window_in_days: i64,
) -> Result<DeleteOutcome, Error> {
tracing::warn!(
"CyberArk Conjur does not support direct secret deletion. Secrets must be removed through policy updates."
);
self.secrets.invalidate(name).await;
Ok(DeleteOutcome::NotSupported)
}
}
impl BaseSecretManager for CyberArkSecretManager {
type Error = Error;
type WriteResponse = ();
type DeleteResponse = DeleteOutcome;
async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
self.async_read_secret(name).await
}
async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
description: Option<&str>,
) -> Result<(), Error> {
self.async_write_secret(name, value, description).await
}
async fn async_delete_secret(
&self,
name: &str,
recovery_window_in_days: i64,
) -> Result<DeleteOutcome, Error> {
self.async_delete_secret(name, recovery_window_in_days)
.await
}
}
fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url {
if !endpoint.path().ends_with('/') {
endpoint.set_path(&format!("{}/", endpoint.path()));
}
endpoint
}

View file

@ -0,0 +1,32 @@
{
"endpoint": "http://conjur.test:8080",
"account": "acct",
"username": "admin",
"api_key": "k3y",
"authenticate_path": "/authn/acct/admin/authenticate",
"token_json": "{\"protected\":\"p\",\"payload\":\"q\",\"signature\":\"s\"}",
"authorization_header": "Token token=\"eyJwcm90ZWN0ZWQiOiJwIiwicGF5bG9hZCI6InEiLCJzaWduYXR1cmUiOiJzIn0=\"",
"policy_path": "/policies/acct/policy/root",
"secrets": [
{
"name": "OPENAI_API_KEY",
"path": "/secrets/acct/variable/OPENAI_API_KEY",
"policy_body": "- !variable \"OPENAI_API_KEY\"\n"
},
{
"name": "team/app/key",
"path": "/secrets/acct/variable/team%2Fapp%2Fkey",
"policy_body": "- !variable \"team/app/key\"\n"
},
{
"name": "a b+c.d-e_f~g",
"path": "/secrets/acct/variable/a%20b%2Bc.d-e_f~g",
"policy_body": "- !variable \"a b+c.d-e_f~g\"\n"
},
{
"name": "needs \"quote\"",
"path": "/secrets/acct/variable/needs%20%22quote%22",
"policy_body": "- !variable \"needs \\\"quote\\\"\"\n"
}
]
}

View file

@ -0,0 +1,516 @@
use std::{sync::Arc, time::Duration};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_secrets_cyberark::{CyberArkSecretManager, DeleteOutcome, Error};
use litellm_secrets_types::SecretValue;
use serde::Deserialize;
use wiremock::{
Match, Mock, MockServer, Request, ResponseTemplate,
matchers::{body_string, header, method, path},
};
const TOKEN_JSON: &str = r#"{"protected":"p","payload":"q","signature":"s"}"#;
#[derive(Deserialize)]
struct ParityFixture {
endpoint: String,
account: String,
username: String,
api_key: String,
authenticate_path: String,
token_json: String,
authorization_header: String,
policy_path: String,
secrets: Vec<ParitySecret>,
}
#[derive(Deserialize)]
struct ParitySecret {
name: String,
path: String,
policy_body: String,
}
#[derive(Debug)]
struct RawPath(String);
impl Match for RawPath {
fn matches(&self, request: &Request) -> bool {
request.url.path() == self.0
}
}
fn fixture() -> ParityFixture {
serde_json::from_str(include_str!("fixtures/parity.json")).unwrap()
}
fn manager(server: &MockServer, ttl: Duration) -> CyberArkSecretManager {
CyberArkSecretManager::with_client(
reqwest::Client::new(),
server.uri().parse().unwrap(),
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(ttl),
)
}
async fn mount_auth(server: &MockServer, expected: u64) {
Mock::given(method("POST"))
.and(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(expected)
.mount(server)
.await;
}
#[tokio::test]
async fn successful_reads_cache_auth_secret_and_redact_values() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
let token = STANDARD.encode(TOKEN_JSON);
Mock::given(path("/secrets/acct/variable/OPENAI_API_KEY"))
.and(header("authorization", format!("Token token=\"{token}\"")))
.respond_with(ResponseTemplate::new(200).set_body_string("sk-live"))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
for _ in 0..2 {
let value = manager
.async_read_secret("OPENAI_API_KEY")
.await
.unwrap()
.unwrap();
assert_eq!(value.expose(), "sk-live");
assert!(!format!("{value:?}").contains("sk-live"));
}
}
#[tokio::test]
async fn concurrent_reads_share_authentication_request() {
let server = MockServer::start().await;
Mock::given(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string(TOKEN_JSON)
.set_delay(Duration::from_millis(20)),
)
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.and(header(
"authorization",
format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)),
))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let (first, second) = tokio::join!(
manager.async_read_secret("key"),
manager.async_read_secret("key")
);
assert_eq!(first.unwrap().unwrap().expose(), "value");
assert_eq!(second.unwrap().unwrap().expose(), "value");
}
#[rstest::rstest]
#[case::not_found(404)]
#[case::unauthorized(401)]
#[case::forbidden(403)]
#[case::server_error(500)]
#[tokio::test]
async fn failed_reads_are_not_cached(#[case] status: u16) {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
let failing = Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(status))
.expect(1)
.mount_as_scoped(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
let result = manager.async_read_secret("key").await;
if status == 404 {
assert_eq!(result.unwrap(), None);
} else {
assert!(matches!(result, Err(Error::Status(actual)) if actual == status));
}
drop(failing);
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("recovered"))
.expect(1)
.mount(&server)
.await;
for _ in 0..2 {
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"recovered"
);
}
}
#[tokio::test]
async fn failed_authentication_is_not_cached_and_does_not_read_secret() {
let server = MockServer::start().await;
let failing = Mock::given(path("/authn/acct/admin/authenticate"))
.respond_with(ResponseTemplate::new(401))
.expect(1)
.mount_as_scoped(&server)
.await;
let unused_secret = Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(0)
.mount_as_scoped(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager.async_read_secret("key").await,
Err(Error::AuthStatus(401))
));
drop(unused_secret);
drop(failing);
mount_auth(&server, 1).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(1)
.mount(&server)
.await;
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[tokio::test]
async fn expired_tokens_and_secrets_are_fetched_again() {
let server = MockServer::start().await;
mount_auth(&server, 2).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_millis(1));
for _ in 0..2 {
assert!(manager.async_read_secret("key").await.unwrap().is_some());
tokio::time::sleep(Duration::from_millis(5)).await;
}
}
#[rstest::rstest]
#[tokio::test]
async fn secret_names_use_python_quote_encoding(
#[values("OPENAI_API_KEY", "team/app/key", "a b+c.d-e_f~g", "needs \"quote\"")] name: &str,
) {
let fixture = fixture();
let secret = fixture
.secrets
.iter()
.find(|secret| secret.name == name)
.unwrap();
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(RawPath(secret.path.clone()))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.expect(1)
.mount(&server)
.await;
assert_eq!(
manager(&server, Duration::from_secs(60))
.async_read_secret(name)
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[rstest::rstest]
#[case(201)]
#[case(409)]
#[case(422)]
#[case(500)]
#[tokio::test]
async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u16) {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/policies/acct/policy/root"))
.and(header("content-type", "application/x-yaml"))
.and(body_string("- !variable \"team/app\"\n"))
.respond_with(ResponseTemplate::new(policy_status))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/team%2Fapp"))
.and(body_string("v"))
.respond_with(ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
manager
.async_write_secret("team/app", &SecretValue::new("v"), None)
.await
.unwrap();
assert_eq!(
manager
.async_read_secret("team/app")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
}
#[tokio::test]
async fn failed_value_write_is_not_cached() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/policies/acct/policy/root"))
.respond_with(ResponseTemplate::new(409))
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.and(body_string("v"))
.respond_with(ResponseTemplate::new(403))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("recovered"))
.expect(1)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager
.async_write_secret("key", &SecretValue::new("v"), None)
.await,
Err(Error::Status(403))
));
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"recovered"
);
}
#[tokio::test]
async fn unsafe_names_fail_before_http_calls() {
let server = MockServer::start().await;
let manager = manager(&server, Duration::from_secs(60));
assert!(matches!(
manager
.async_write_secret("../etc", &SecretValue::new("v"), None)
.await,
Err(Error::Operation(
litellm_secrets_types::Error::UnsafeSecretName
))
));
}
#[tokio::test]
async fn delete_invalidates_cache_and_reports_not_supported() {
let server = MockServer::start().await;
mount_auth(&server, 1).await;
Mock::given(path("/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("v"))
.expect(2)
.mount(&server)
.await;
let manager = manager(&server, Duration::from_secs(60));
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
assert_eq!(
manager.async_delete_secret("key", 7).await.unwrap(),
DeleteOutcome::NotSupported
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"v"
);
}
#[test]
fn new_validates_credentials_before_license_and_configuration() {
let empty: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
Arc::new(|_: &str| None);
assert!(matches!(
CyberArkSecretManager::new(empty, true),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| (name == "CYBERARK_API_KEY").then(|| "k3y".into())),
false
),
Err(Error::EnterpriseRequired)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| (name == "CYBERARK_CLIENT_CERT").then(|| "cert".into())),
true
),
Err(Error::MissingCredentials)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_REFRESH_INTERVAL" => Some("abc".into()),
_ => None,
}),
true
),
Err(Error::RefreshInterval)
));
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_API_KEY" => Some("k3y".into()),
"CYBERARK_API_BASE" => Some("not a url".into()),
_ => None,
}),
true
),
Err(Error::Endpoint)
));
}
#[tokio::test]
async fn new_reads_environment_defaults_end_to_end() {
let server = MockServer::start().await;
Mock::given(path("/authn/default/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.mount(&server)
.await;
Mock::given(path("/secrets/default/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let endpoint = server.uri();
let manager = CyberArkSecretManager::new(
Arc::new(move |name: &str| match name {
"CYBERARK_API_BASE" => Some(endpoint.clone()),
"CYBERARK_API_KEY" => Some("k3y".into()),
_ => None,
}),
true,
)
.unwrap();
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[test]
fn new_reports_missing_client_certificate_files() {
assert!(matches!(
CyberArkSecretManager::new(
Arc::new(|name: &str| match name {
"CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()),
"CYBERARK_CLIENT_KEY" => Some("/missing/key".into()),
_ => None,
}),
true
),
Err(Error::ClientCertificate)
));
}
#[tokio::test]
async fn trailing_slash_endpoint_preserves_base_path() {
let server = MockServer::start().await;
Mock::given(path("/prefix/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON))
.expect(1)
.mount(&server)
.await;
Mock::given(path("/prefix/secrets/acct/variable/key"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let endpoint = format!("{}/prefix/", server.uri()).parse().unwrap();
let manager = CyberArkSecretManager::with_client(
reqwest::Client::new(),
endpoint,
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(Duration::from_secs(60)),
);
assert_eq!(
manager
.async_read_secret("key")
.await
.unwrap()
.unwrap()
.expose(),
"value"
);
}
#[test]
fn parity_fixture_matches_authentication_contract() {
let fixture = fixture();
assert_eq!(fixture.endpoint, "http://conjur.test:8080");
assert_eq!(fixture.account, "acct");
assert_eq!(fixture.username, "admin");
assert_eq!(fixture.api_key, "k3y");
assert_eq!(fixture.authenticate_path, "/authn/acct/admin/authenticate");
assert_eq!(fixture.token_json, TOKEN_JSON);
assert_eq!(
fixture.authorization_header,
format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON))
);
assert_eq!(fixture.policy_path, "/policies/acct/policy/root");
assert_eq!(fixture.secrets.len(), 4);
assert_eq!(
fixture.secrets[1].policy_body,
"- !variable \"team/app/key\"\n"
);
}

View file

@ -9,11 +9,13 @@ repository.workspace = true
default = []
aws = ["dep:litellm-secrets-aws"]
google = ["dep:litellm-secrets-google"]
cyberark = ["dep:litellm-secrets-cyberark"]
[dependencies]
litellm-secrets-types.workspace = true
litellm-secrets-aws = { workspace = true, optional = true }
litellm-secrets-google = { workspace = true, optional = true }
litellm-secrets-cyberark = { workspace = true, optional = true }
litellm-core-utils.workspace = true
base64.workspace = true
serde.workspace = true

View file

@ -30,4 +30,7 @@ pub enum Error {
#[cfg(feature = "google")]
#[error(transparent)]
Google(#[from] litellm_secrets_google::Error),
#[cfg(feature = "cyberark")]
#[error(transparent)]
Cyberark(#[from] litellm_secrets_cyberark::Error),
}

View file

@ -13,6 +13,8 @@ pub enum SecretManager {
GoogleKms(crate::google::GoogleKms),
#[cfg(feature = "google")]
GoogleSecretManager(crate::google::GoogleSecretManager),
#[cfg(feature = "cyberark")]
Cyberark(crate::cyberark::CyberArkSecretManager),
}
impl SecretManager {
@ -27,6 +29,8 @@ impl SecretManager {
Self::GoogleKms(_) => KeyManagementSystem::GoogleKms,
#[cfg(feature = "google")]
Self::GoogleSecretManager(_) => KeyManagementSystem::GoogleSecretManager,
#[cfg(feature = "cyberark")]
Self::Cyberark(_) => KeyManagementSystem::Cyberark,
}
}
}
@ -78,6 +82,12 @@ pub async fn get_secret_from_manager(
.get_secret_from_google_secret_manager(secret_name)
.await
.map_err(Error::from),
#[cfg(feature = "cyberark")]
SecretManager::Cyberark(client) => client
.async_read_secret(secret_name)
.await
.map(|value| value.map(Secret::String))
.map_err(Error::from),
}
}

View file

@ -17,5 +17,7 @@ pub use state::{SecretManagerState, secret_manager_would_be_consulted};
#[cfg(feature = "aws")]
pub use litellm_secrets_aws as aws;
#[cfg(feature = "cyberark")]
pub use litellm_secrets_cyberark as cyberark;
#[cfg(feature = "google")]
pub use litellm_secrets_google as google;

View file

@ -105,3 +105,56 @@ async fn google_handler_requires_canonical_base64_and_preserves_plaintext_whites
Err(Error::MissingCiphertext)
));
}
#[cfg(feature = "cyberark")]
#[tokio::test]
async fn cyberark_handler_reads_values_and_surfaces_errors() {
use std::time::Duration;
use litellm_secrets::{
Error, KeyManagementSettings, SecretManager, SecretValue, cyberark::CyberArkSecretManager,
get_secret_from_manager,
};
use wiremock::{
Mock, MockServer, ResponseTemplate,
matchers::{body_string, path},
};
let server = MockServer::start().await;
Mock::given(path("/authn/acct/admin/authenticate"))
.and(body_string("k3y"))
.respond_with(ResponseTemplate::new(200).set_body_string("token"))
.mount(&server)
.await;
Mock::given(path("/secrets/acct/variable/KEY"))
.respond_with(ResponseTemplate::new(200).set_body_string("value"))
.mount(&server)
.await;
let manager = SecretManager::Cyberark(CyberArkSecretManager::with_client(
reqwest::Client::new(),
server.uri().parse().unwrap(),
"acct".into(),
"admin".into(),
SecretValue::new("k3y"),
Some(Duration::from_secs(60)),
));
assert_eq!(
manager.system(),
litellm_secrets::KeyManagementSystem::Cyberark
);
let settings = KeyManagementSettings::default();
let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None)
.await
.unwrap()
.unwrap();
assert_eq!(value.as_str(), Some("value"));
Mock::given(path("/secrets/acct/variable/ERROR"))
.respond_with(ResponseTemplate::new(500))
.mount(&server)
.await;
assert!(matches!(
get_secret_from_manager(&manager, "ERROR", &settings, &|_: &str| None).await,
Err(Error::Cyberark(_))
));
}

View file

@ -689,6 +689,7 @@ recraft_models: Set = set()
cometapi_models: Set = set()
oci_models: Set = set()
vercel_ai_gateway_models: Set = set()
edenai_models: Set = set() # mutable-ok: filled from the price map at import, like the sibling provider sets
volcengine_models: Set = set()
wandb_models: Set = set(WANDB_MODELS)
ovhcloud_models: Set = set()
@ -763,6 +764,8 @@ def _populate_provider_model_sets(model_cost_map: Dict) -> None:
openrouter_models.add(key)
elif value.get("litellm_provider") == "vercel_ai_gateway":
vercel_ai_gateway_models.add(key)
elif value.get("litellm_provider") == "edenai":
edenai_models.add(key)
elif value.get("litellm_provider") == "datarobot":
datarobot_models.add(key)
elif value.get("litellm_provider") == "vertex_ai-text-models":
@ -1111,6 +1114,7 @@ model_list = list(
| oci_models
| heroku_models
| vercel_ai_gateway_models
| edenai_models
| volcengine_models
| wandb_models
| ovhcloud_models
@ -1139,6 +1143,7 @@ def _build_models_by_provider() -> dict:
"baseten": baseten_models,
"openrouter": openrouter_models,
"vercel_ai_gateway": vercel_ai_gateway_models,
"edenai": edenai_models,
"datarobot": datarobot_models,
"vertex_ai": vertex_chat_models
| vertex_text_models
@ -2117,6 +2122,30 @@ if TYPE_CHECKING:
from .llms.vercel_ai_gateway.chat.transformation import (
VercelAIGatewayConfig as VercelAIGatewayConfig,
)
from .llms.edenai.chat.transformation import (
EdenAIChatConfig as EdenAIChatConfig,
)
from .llms.edenai.responses.transformation import (
EdenAIResponsesAPIConfig as EdenAIResponsesAPIConfig,
)
from .llms.edenai.messages.transformation import (
EdenAIAnthropicMessagesConfig as EdenAIAnthropicMessagesConfig,
)
from .llms.edenai.embedding.transformation import (
EdenAIEmbeddingConfig as EdenAIEmbeddingConfig,
)
from .llms.edenai.audio_transcription.transformation import (
EdenAIAudioTranscriptionConfig as EdenAIAudioTranscriptionConfig,
)
from .llms.edenai.text_to_speech.transformation import (
EdenAITextToSpeechConfig as EdenAITextToSpeechConfig,
)
from .llms.edenai.image_generation.transformation import (
EdenAIImageGenerationConfig as EdenAIImageGenerationConfig,
)
from .llms.edenai.videos.transformation import (
EdenAIVideoConfig as EdenAIVideoConfig,
)
from .llms.ovhcloud.chat.transformation import (
OVHCloudChatConfig as OVHCloudChatConfig,
)

View file

@ -327,6 +327,14 @@ LLM_CONFIG_NAMES: Final = (
"InceptionChatConfig",
"HyperbolicChatConfig",
"VercelAIGatewayConfig",
"EdenAIChatConfig",
"EdenAIResponsesAPIConfig",
"EdenAIAnthropicMessagesConfig",
"EdenAIEmbeddingConfig",
"EdenAIAudioTranscriptionConfig",
"EdenAITextToSpeechConfig",
"EdenAIImageGenerationConfig",
"EdenAIVideoConfig",
"OVHCloudChatConfig",
"OVHCloudEmbeddingConfig",
"CometAPIEmbeddingConfig",
@ -1232,6 +1240,17 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
".llms.vercel_ai_gateway.chat.transformation",
"VercelAIGatewayConfig",
),
"EdenAIChatConfig": (".llms.edenai.chat.transformation", "EdenAIChatConfig"),
"EdenAIResponsesAPIConfig": (".llms.edenai.responses.transformation", "EdenAIResponsesAPIConfig"),
"EdenAIAnthropicMessagesConfig": (".llms.edenai.messages.transformation", "EdenAIAnthropicMessagesConfig"),
"EdenAIEmbeddingConfig": (".llms.edenai.embedding.transformation", "EdenAIEmbeddingConfig"),
"EdenAIAudioTranscriptionConfig": (
".llms.edenai.audio_transcription.transformation",
"EdenAIAudioTranscriptionConfig",
),
"EdenAITextToSpeechConfig": (".llms.edenai.text_to_speech.transformation", "EdenAITextToSpeechConfig"),
"EdenAIImageGenerationConfig": (".llms.edenai.image_generation.transformation", "EdenAIImageGenerationConfig"),
"EdenAIVideoConfig": (".llms.edenai.videos.transformation", "EdenAIVideoConfig"),
"OVHCloudChatConfig": (".llms.ovhcloud.chat.transformation", "OVHCloudChatConfig"),
"OVHCloudEmbeddingConfig": (
".llms.ovhcloud.embedding.transformation",

View file

@ -11,6 +11,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": "dangerous-tool-use-2026-09-03",
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": "fast-mode-2026-02-01",
"files-api-2025-04-14": "files-api-2025-04-14",
@ -44,6 +45,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": null,
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": "files-api-2025-04-14",
@ -76,6 +78,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": null,
"dangerous-tool-use-2026-09-03": null,
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
@ -109,6 +112,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": "dangerous-tool-use-2026-09-03",
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
@ -143,6 +147,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": "dangerous-tool-use-2026-09-03",
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
@ -177,6 +182,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": "dangerous-tool-use-2026-09-03",
"effort-2025-11-24": null,
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
@ -210,6 +216,7 @@
"computer-use-2025-11-24": "computer-use-2025-11-24",
"context-1m-2025-08-07": "context-1m-2025-08-07",
"context-management-2025-06-27": "context-management-2025-06-27",
"dangerous-tool-use-2026-09-03": null,
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": "fast-mode-2026-02-01",
"files-api-2025-04-14": "files-api-2025-04-14",

View file

@ -80,6 +80,8 @@ class _AsyncRedisCommands(Protocol):
def ttl(self, name: str) -> Awaitable[int]: ...
def expire(self, name: str, time: int) -> Awaitable[bool]: ...
def rpush(self, name: str, *values: str | bytes | float) -> Awaitable[int]: ...
def lpop(self, name: str, count: int | None = None) -> Awaitable[object]: ...
@ -1948,6 +1950,14 @@ class RedisCache(BaseCache):
_record_swallowed_redis_failure(self._circuit_breaker, e)
return None
@_redis_circuit_breaker_guard
async def async_refresh_ttl(self, key: str, ttl: int | None = None) -> bool:
"""EXPIRE an existing key without touching its value. False when the key is absent."""
_used_ttl: Final = self.get_ttl(ttl=ttl)
if _used_ttl is None:
return False
return await self._async_commands().expire(self.check_and_fix_namespace(key=key), _used_ttl)
@_redis_circuit_breaker_guard
async def async_rpush(
self,

View file

@ -750,6 +750,7 @@ LITELLM_CHAT_PROVIDERS: Final = [
"inception",
"vercel_ai_gateway",
"wandb",
"edenai",
"ovhcloud",
"lemonade",
"docker_model_runner",
@ -925,6 +926,7 @@ openai_compatible_endpoints: Final[list] = [
"https://api.hyperbolic.xyz/v1",
"https://ai-gateway.helicone.ai/",
"https://ai-gateway.vercel.sh/v1",
"https://api.edenai.run/v3",
"https://api.inference.wandb.ai/v1",
"https://api.clarifai.com/v2/ext/openai/v1",
"https://api.libertai.io/v1",
@ -994,6 +996,7 @@ openai_compatible_providers: Final[list] = [
"hyperbolic",
"vercel_ai_gateway",
"aiml",
"edenai",
"wandb",
"cometapi",
"clarifai",

View file

@ -388,6 +388,7 @@ def image_generation(
litellm.LlmProviders.DASHSCOPE,
litellm.LlmProviders.QWENCLOUD,
litellm.LlmProviders.QWEN_AI_PLATFORM,
litellm.LlmProviders.EDENAI,
):
if image_generation_config is None:
raise ValueError(f"image generation config is not supported for {custom_llm_provider}")

View file

@ -1,14 +1,18 @@
"""
Handles Batching + sending Httpx Post requests to slack
Slack alerts are sent every 10s or when events are greater than X events
Slack alerts are sent every DEFAULT_FLUSH_INTERVAL_SECONDS or when events are greater than X events
see custom_batch_logger.py for more details / defaults
"""
from collections import Counter
from collections.abc import Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_proxy_logger
from litellm.types.integrations.slack_alerting import AlertQueueItem, AlertType
from .ms_teams import MS_TEAMS_ALERTING_DESTINATION, build_ms_teams_payload
@ -20,26 +24,20 @@ else:
SlackAlertingType = Any
def squash_payloads(queue):
squashed: Final = {}
if len(queue) == 0:
return squashed
if len(queue) == 1:
return {"key": {"item": queue[0], "count": 1}}
@dataclass(frozen=True, slots=True)
class SquashedAlert:
item: AlertQueueItem
count: int
for item in queue:
url = item["url"]
alert_type = item["alert_type"]
_key = (url, alert_type)
if _key in squashed:
squashed[_key]["count"] += 1
# Merge the payloads
def _squash_key(item: AlertQueueItem) -> tuple[str, AlertType | str, str]:
return (item["url"], item["alert_type"], item["payload"]["text"])
else:
squashed[_key] = {"item": item, "count": 1}
return squashed
def squash_payloads(queue: Sequence[AlertQueueItem]) -> tuple[SquashedAlert, ...]:
counts: Final = Counter(_squash_key(item) for item in queue)
first_item_by_key: Final = {_squash_key(item): item for item in reversed(queue)}
return tuple(SquashedAlert(item=first_item_by_key[key], count=count) for key, count in counts.items())
def _print_alerting_payload_warning(payload: dict, slackAlertingInstance: SlackAlertingType):
@ -53,17 +51,15 @@ def _print_alerting_payload_warning(payload: dict, slackAlertingInstance: SlackA
verbose_proxy_logger.warning(payload)
async def send_to_webhook(slackAlertingInstance: SlackAlertingType, item, count):
async def send_to_webhook(slackAlertingInstance: SlackAlertingType, item: AlertQueueItem, count: int) -> None:
"""
Send a single slack alert to the webhook
"""
import json
payload: Final = item.get("payload", {})
text: Final = item["payload"]["text"]
payload: Final = {"text": text if count == 1 else f"[Num Alerts: {count}]\n\n{text}"}
try:
if count > 1:
payload["text"] = f"[Num Alerts: {count}]\n\n{payload['text']}"
request_body: Final = (
build_ms_teams_payload(payload["text"]) if item.get("format") == MS_TEAMS_ALERTING_DESTINATION else payload
)

View file

@ -33,6 +33,7 @@ from litellm.litellm_core_utils.exception_mapping_utils import (
_add_key_name_and_team_to_alert,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
@ -99,6 +100,7 @@ class SlackAlerting(CustomBatchLogger):
alerting_args={},
default_webhook_url: str | None = None,
alert_type_config: dict[str, dict] | None = None,
async_http_handler: AsyncHTTPHandler | None = None,
**kwargs,
):
if alerting_threshold is None:
@ -107,7 +109,9 @@ class SlackAlerting(CustomBatchLogger):
self.alerting = alerting
self.alert_types = alert_types
self.internal_usage_cache = internal_usage_cache or DualCache()
self.async_http_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
self.async_http_handler = async_http_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
self.alert_to_webhook_url = process_slack_alerting_variables(alert_to_webhook_url=alert_to_webhook_url)
self.is_running = False
self.alerting_args = SlackAlertingArgs(**alerting_args)
@ -1583,12 +1587,12 @@ Model Info:
if not self.log_queue:
return
squashed_queue: Final = squash_payloads(self.log_queue)
tasks: Final = [
send_to_webhook(slackAlertingInstance=self, item=item["item"], count=item["count"])
for item in squashed_queue.values()
]
await asyncio.gather(*tasks)
await asyncio.gather(
*(
send_to_webhook(slackAlertingInstance=self, item=squashed.item, count=squashed.count)
for squashed in squash_payloads(self.log_queue)
)
)
self.log_queue.clear()
async def _flush_digest_buckets(self):

View file

@ -427,7 +427,7 @@ class LLMCallSpanData:
# plain ``.get`` — no repeated ``isinstance`` guards.
raw_response: Final = payload.get("response")
response: Final = cast(Mapping[str, object], raw_response if isinstance(raw_response, dict) else {})
choices_out: Final = _dicts(response.get("choices")) or _responses_choices(response)
choices_out: Final = _dicts(response.get("choices")) or _responses_choices(response) or _ocr_choices(response)
# ``finish_reasons`` is metadata, not content, so derive it from
# ``choices_out`` before gating. The raw message/choice bodies are only
# retained when content capture is enabled (see ``capture_span_content``);
@ -752,6 +752,22 @@ def _responses_choices(response: Mapping[str, object]) -> tuple[_Choice, ...]:
return (choice,)
def _ocr_choices(response: Mapping[str, object]) -> tuple[_Choice, ...]:
markdowns: Final = tuple(
text for page in _dicts(response.get("pages")) if (text := as_str(page.get("markdown"))) is not None
)
if not markdowns:
return ()
message: Final[_AssistantMessage] = {
"role": "assistant",
"content": "\n\n".join(markdowns),
"refusal": None,
"tool_calls": None,
}
choice: Final[_Choice] = {"message": message, "finish_reason": None}
return (choice,)
def _responses_parts_text(parts: tuple[Mapping[str, object], ...], part_type: str, field: str) -> str | None:
texts: Final = tuple(
text for part in parts if part.get("type") == part_type if (text := as_str(part.get(field))) is not None

View file

@ -4,7 +4,8 @@ import copy
import logging
import re
from collections.abc import Iterable, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol
import httpx
from pydantic import TypeAdapter, ValidationError
@ -703,3 +704,24 @@ def redact_nested_match_and_regex_keys(
except Exception:
return payload
return redacted
RESPONSE_COST_HEADER: Final = "llm_provider-x-litellm-response-cost"
_NO_HEADERS: Final[Mapping[str, object]] = MappingProxyType({})
class _CarriesHiddenParams(Protocol):
_hidden_params: dict[str, object] # mutable-ok: the responses billed here keep hidden params in a plain dict
def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: float | None) -> None:
"""Record a provider-reported cost where the cost calculator looks before the price map."""
if cost is None:
return
hidden_params: Final = response._hidden_params # pyright: ignore[reportPrivateUsage] # no public accessor
additional_headers: Final[object] = hidden_params.get("additional_headers")
merged: Final[dict[str, object]] = { # mutable-ok: assigned into the plain-dict hidden params
**(additional_headers if isinstance(additional_headers, Mapping) else _NO_HEADERS),
RESPONSE_COST_HEADER: cost,
}
hidden_params["additional_headers"] = merged # rebind-ok: the caller's record is the point

View file

@ -362,6 +362,9 @@ def get_llm_provider(
elif endpoint == "https://ai-gateway.vercel.sh/v1":
custom_llm_provider = "vercel_ai_gateway"
dynamic_api_key = get_secret_str("VERCEL_AI_GATEWAY_API_KEY")
elif endpoint == "https://api.edenai.run/v3":
custom_llm_provider = "edenai" # rebind-ok: api_base detection resolves the provider in place
dynamic_api_key = get_secret_str("EDENAI_API_KEY")
elif endpoint == "https://api.inference.wandb.ai/v1":
custom_llm_provider = "wandb"
dynamic_api_key = get_secret_str("WANDB_API_KEY")
@ -853,6 +856,9 @@ def _get_openai_compatible_provider_info(
api_base,
dynamic_api_key,
) = litellm.VercelAIGatewayConfig()._get_openai_compatible_provider_info(api_base, api_key)
elif custom_llm_provider == "edenai":
api_base = litellm.EdenAIChatConfig.get_api_base(api_base) # rebind-ok: chain resolves in place
dynamic_api_key = litellm.EdenAIChatConfig.get_api_key(api_key) # rebind-ok: chain resolves in place
elif custom_llm_provider == "aiml":
(
api_base,

View file

@ -69,7 +69,11 @@ from litellm.litellm_core_utils.classifier_logging import (
classifier_input_snapshot,
is_classifier_call,
)
from litellm.litellm_core_utils.core_helpers import is_expected_client_error, reconstruct_model_name
from litellm.litellm_core_utils.core_helpers import (
is_expected_client_error,
reconstruct_model_name,
set_response_cost_in_hidden_params,
)
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.internal_call_metadata import (
MODEL_ACCESS_GROUP_METADATA_KEY,
@ -3918,6 +3922,7 @@ class Logging(LiteLLMLoggingBaseClass):
):
## return unified Usage object
if isinstance(result.response.usage, ResponseAPIUsage):
set_response_cost_in_hidden_params(result.response, result.response.usage.cost)
transformed_usage: Final = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
result.response.usage
)

View file

@ -65,7 +65,7 @@ def _build_secret_patterns() -> "re.Pattern[str]":
# private_key with PEM-aware value capture
r"""private_key['\"]?\s*[:=]\s*['\"]?(?:-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----|[^\s,'\"})\]{}>]+)""",
r"(?:master_key|xai_key|database_url|db_url|connection_string|"
r"aws_secret_access_key|aws_session_token|aws_access_key_id|"
r"aws_secret_access_key|aws_session_token|aws_access_key_id|s3_secret_access_key|s3_access_key_id|"
r"signing_key|encryption_key|"
r"auth_token|access_token|refresh_token|"
r"slack_webhook_url|webhook_url|"

View file

@ -272,6 +272,19 @@ class BaseVideoConfig(ABC):
) -> VideoObject:
pass
async def async_transform_video_status_retrieve_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str | None = None,
) -> VideoObject:
"""Async transform video status retrieve response."""
return self.transform_video_status_retrieve_response(
raw_response=raw_response,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
)
def transform_video_create_character_request(
self,
name: str,

View file

@ -33,6 +33,7 @@ from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
CommonBatchFilesUtils,
merge_bedrock_aws_request_params,
resolve_s3_bucket_owner,
resolve_s3_encryption_key_id,
)
@ -51,6 +52,26 @@ _S3_BATCH_FILE_UUID_SUFFIX_PATTERN: Final = re.compile(
_BEDROCK_TAGS_ADAPTER: Final[TypeAdapter[list[BedrockTag]]] = TypeAdapter(list[BedrockTag])
def _build_s3_input_config(s3_uri: str, s3_bucket_owner: str | None) -> BedrockS3InputDataConfig:
if s3_bucket_owner is None:
return BedrockS3InputDataConfig(s3Uri=s3_uri)
return BedrockS3InputDataConfig(s3Uri=s3_uri, s3BucketOwner=s3_bucket_owner)
def _build_s3_output_config(
s3_uri: str, s3_bucket_owner: str | None, s3_encryption_key_id: str | None
) -> BedrockS3OutputDataConfig:
if s3_bucket_owner is None:
if s3_encryption_key_id is None:
return BedrockS3OutputDataConfig(s3Uri=s3_uri)
return BedrockS3OutputDataConfig(s3Uri=s3_uri, s3EncryptionKeyId=s3_encryption_key_id)
if s3_encryption_key_id is None:
return BedrockS3OutputDataConfig(s3Uri=s3_uri, s3BucketOwner=s3_bucket_owner)
return BedrockS3OutputDataConfig(
s3Uri=s3_uri, s3BucketOwner=s3_bucket_owner, s3EncryptionKeyId=s3_encryption_key_id
)
def _validate_bedrock_tags(raw_tags: object) -> list[BedrockTag]:
try:
return _BEDROCK_TAGS_ADAPTER.validate_python(raw_tags, strict=True)
@ -214,25 +235,23 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
job_name: Final = self.common_utils.generate_unique_job_name(model, prefix="litellm")
output_key: Final = f"litellm-batch-outputs/{job_name}/"
# Build input data config
input_data_config: Final[BedrockInputDataConfig] = {
"s3InputDataConfig": BedrockS3InputDataConfig(s3Uri=f"s3://{input_bucket}/{input_key}")
}
# Build output data config
s3_output_config: Final[BedrockS3OutputDataConfig] = BedrockS3OutputDataConfig(
s3Uri=f"s3://{output_bucket}/{output_key}"
)
# Add optional KMS encryption key ID if provided
s3_encryption_key_id = resolve_s3_encryption_key_id(
s3_bucket_owner: Final = resolve_s3_bucket_owner(litellm_params=litellm_params, optional_params=optional_params)
s3_encryption_key_id: Final = resolve_s3_encryption_key_id(
litellm_params=litellm_params,
optional_params=optional_params,
)
if s3_encryption_key_id:
s3_output_config["s3EncryptionKeyId"] = s3_encryption_key_id
output_data_config: Final[BedrockOutputDataConfig] = {"s3OutputDataConfig": s3_output_config}
input_data_config: Final[BedrockInputDataConfig] = {
"s3InputDataConfig": _build_s3_input_config(
s3_uri=f"s3://{input_bucket}/{input_key}", s3_bucket_owner=s3_bucket_owner
)
}
output_data_config: Final[BedrockOutputDataConfig] = {
"s3OutputDataConfig": _build_s3_output_config(
s3_uri=f"s3://{output_bucket}/{output_key}",
s3_bucket_owner=s3_bucket_owner,
s3_encryption_key_id=s3_encryption_key_id,
)
}
# Create Bedrock batch request with proper typing
bedrock_request: Final[BedrockCreateBatchRequest] = {

View file

@ -1555,11 +1555,33 @@ def resolve_s3_encryption_key_id(
Precedence: `s3_encryption_key_id` in litellm_params, then optional_params
(client-side / request params), then the AWS_S3_ENCRYPTION_KEY_ID env var.
"""
return _resolve_s3_setting("s3_encryption_key_id", "AWS_S3_ENCRYPTION_KEY_ID", litellm_params, optional_params)
def resolve_s3_bucket_owner(
litellm_params: Mapping[str, object],
optional_params: Mapping[str, object] | None = None,
) -> str | None:
"""
Resolve the AWS account id that owns the S3 buckets used by Bedrock batch jobs.
Precedence: `s3_bucket_owner` in litellm_params, then optional_params
(client-side / request params), then the AWS_S3_BUCKET_OWNER env var.
"""
return _resolve_s3_setting("s3_bucket_owner", "AWS_S3_BUCKET_OWNER", litellm_params, optional_params)
def _resolve_s3_setting(
param_name: str,
env_var: str,
litellm_params: Mapping[str, object],
optional_params: Mapping[str, object] | None,
) -> str | None:
candidates: Final = tuple(
source.get("s3_encryption_key_id") for source in (litellm_params, optional_params) if source is not None
source.get(param_name) for source in (litellm_params, optional_params) if source is not None
)
explicit: Final = next((value for value in candidates if isinstance(value, str) and value), None)
return explicit or get_secret_str("AWS_S3_ENCRYPTION_KEY_ID")
return explicit or get_secret_str(env_var)
class CommonBatchFilesUtils:

View file

@ -533,6 +533,9 @@ class AmazonAnthropicClaudeMessagesConfig(
if anthropic_model_info.is_eager_input_streaming_used(tools):
beta_set.add(ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER)
if anthropic_messages_optional_request_params.get("safeguards") is not None:
beta_set.add(ANTHROPIC_BETA_HEADER_VALUES.DANGEROUS_TOOL_USE_2026_09_03.value)
self._filter_context_management_for_bedrock_invoke(
anthropic_messages_request=anthropic_messages_request,
beta_set=beta_set,

View file

@ -8881,7 +8881,7 @@ class BaseLLMHTTPHandler:
url=url,
headers=headers,
)
return video_status_provider_config.transform_video_status_retrieve_response(
return await video_status_provider_config.async_transform_video_status_retrieve_response(
raw_response=response,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,

View file

@ -103,7 +103,7 @@ def missing_dashscope_family_key_message(custom_llm_provider: str) -> str:
)
if custom_llm_provider == "qwen_ai_platform":
return (
"Missing API key for Qwen AI Platform. Set QWEN_AI_PLATFORM_API_KEY or "
"Missing API key for Qianwen AI Platform. Set QWEN_AI_PLATFORM_API_KEY or "
"DASHSCOPE_API_KEY environment variable or pass api_key parameter."
)
return "Missing API key for DashScope. Set DASHSCOPE_API_KEY environment variable or pass api_key parameter."

View file

@ -23,7 +23,7 @@ def _require_qwen_ai_platform_api_key(api_key: str | None) -> str:
resolved: Final = _resolve_qwen_ai_platform_api_key(api_key)
if resolved is None:
raise ValueError(
"Qwen AI Platform API key is required. Set 'QWEN_AI_PLATFORM_API_KEY' or 'DASHSCOPE_API_KEY' env var "
"Qianwen AI Platform API key is required. Set 'QWEN_AI_PLATFORM_API_KEY' or 'DASHSCOPE_API_KEY' env var "
"or pass api_key explicitly."
)
return resolved

Some files were not shown because too many files have changed in this diff Show more