From 9b747bf337fb0423701dcb53845646ff08fd4288 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Thu, 24 Sep 2026 09:19:14 -0700 Subject: [PATCH] feat(caching): give RedisCache a pub/sub and lease surface and drop raw client reaches Add async_publish and async_subscribe to the Python RedisCache, returning a RedisSubscription that yields RedisMessage values and skips redis-py's subscribe acks, plus connection_pool_status for the debug endpoint. Cluster clients raise NotImplementedError, matching the existing skip behaviour. Move the config sync and auth cache invalidation modules onto that surface and rewrite the MCP refresh lock onto async_register_script, so no proxy code reaches into init_async_client or redis_client any more. This defines the surface a native RedisCache has to match. On the Rust side add the PubSubCache and MessageStream capability traits and implement them in cache-redis with a dedicated subscription connection driven by a reader thread, tested against an in-process RESP server over a real socket. --- litellm-rust/crates/cache-redis/Cargo.toml | 2 +- .../crates/cache-redis/src/connection.rs | 34 ++- litellm-rust/crates/cache-redis/src/lib.rs | 2 + litellm-rust/crates/cache-redis/src/pubsub.rs | 154 +++++++++++ .../crates/cache-redis/tests/pubsub.rs | 105 ++++++++ .../crates/cache-redis/tests/support/mod.rs | 2 + .../cache-redis/tests/support/server.rs | 253 ++++++++++++++++++ litellm-rust/crates/cache/src/capabilities.rs | 35 +++ litellm-rust/crates/cache/src/lib.rs | 5 +- litellm/caching/redis_cache.py | 90 ++++++- .../redis_distributed_lock.py | 84 +++--- .../runtime_refresh_coordinator.py | 9 +- .../auth_cache_invalidation_pubsub.py | 59 ++-- .../proxy/common_utils/config_sync_pubsub.py | 86 ++---- litellm/proxy/common_utils/debug_utils.py | 23 +- .../test_litellm/caching/test_redis_cache.py | 243 ++++++++++------- .../test_redis_distributed_lock.py | 211 ++++++++------- .../test_auth_cache_invalidation_pubsub.py | 143 ++++++---- .../common_utils/test_config_sync_pubsub.py | 233 ++++++++-------- .../proxy/common_utils/test_debug_utils.py | 25 ++ 20 files changed, 1251 insertions(+), 547 deletions(-) create mode 100644 litellm-rust/crates/cache-redis/src/pubsub.rs create mode 100644 litellm-rust/crates/cache-redis/tests/pubsub.rs create mode 100644 litellm-rust/crates/cache-redis/tests/support/server.rs diff --git a/litellm-rust/crates/cache-redis/Cargo.toml b/litellm-rust/crates/cache-redis/Cargo.toml index a234286a338..46f2e438f6a 100644 --- a/litellm-rust/crates/cache-redis/Cargo.toml +++ b/litellm-rust/crates/cache-redis/Cargo.toml @@ -9,7 +9,7 @@ repository.workspace = true litellm-cache.workspace = true redis = { version = "1.7.0", features = ["cluster", "tls-rustls"] } r2d2 = "0.8.10" -tokio.workspace = true +tokio = { workspace = true, features = ["sync"] } [dev-dependencies] litellm-cache-testing.workspace = true diff --git a/litellm-rust/crates/cache-redis/src/connection.rs b/litellm-rust/crates/cache-redis/src/connection.rs index 2f58e2a9b80..f8ec23c03f4 100644 --- a/litellm-rust/crates/cache-redis/src/connection.rs +++ b/litellm-rust/crates/cache-redis/src/connection.rs @@ -21,7 +21,10 @@ const REDIS_POOL_SIZE: u32 = 16; #[allow(private_interfaces)] pub enum Connections { - Pool(r2d2::Pool), + Pool { + pool: r2d2::Pool, + client: Box, + }, Cluster(r2d2::Pool), Fixed(Mutex), } @@ -35,7 +38,7 @@ where operation: impl FnOnce(&mut ConnectionRef<'_>) -> Result, ) -> Result { match self { - Self::Pool(pool) => { + 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)); @@ -70,7 +73,13 @@ where pub fn open(url: &str, topology: &RedisTopology) -> Result { match topology { - RedisTopology::Standalone => Ok(Self::Pool(pool(ConnectionManager::open(url)?)?)), + RedisTopology::Standalone => { + let client = redis::Client::open(url).map_err(|_| Error::Unavailable)?; + Ok(Self::Pool { + pool: pool(ConnectionManager(client.clone()))?, + client: Box::new(client), + }) + } RedisTopology::Cluster { startup_nodes } => Ok(Self::Cluster(pool( ClusterConnectionManager::open(url, startup_nodes)?, )?)), @@ -81,11 +90,20 @@ where /// checked out right now return to the pool, and a caller-owned connection stays open. pub fn disconnect(&self) { match self { - Self::Pool(pool) => close_idle(pool), + Self::Pool { pool, .. } => close_idle(pool), Self::Cluster(pool) => close_idle(pool), Self::Fixed(_) => {} } } + + /// A connection outside the pool for a subscription to take over; a subscribed RESP2 + /// connection accepts nothing else, so it can never go back to the pool. + pub(crate) fn subscription_connection(&self) -> Result { + match self { + Self::Pool { client, .. } => client.get_connection().map_err(|_| Error::Unavailable), + Self::Cluster(_) | Self::Fixed(_) => Err(Error::UnsupportedOperation), + } + } } fn pool(manager: M) -> Result, Error> { @@ -119,14 +137,6 @@ pub(crate) struct PooledConnection { /// open, so any connection whose operation failed is discarded instead of being reused. pub(crate) struct ConnectionManager(redis::Client); -impl ConnectionManager { - fn open(url: &str) -> Result { - redis::Client::open(url) - .map(Self) - .map_err(|_| Error::Unavailable) - } -} - impl r2d2::ManageConnection for ConnectionManager { type Connection = PooledConnection; type Error = redis::RedisError; diff --git a/litellm-rust/crates/cache-redis/src/lib.rs b/litellm-rust/crates/cache-redis/src/lib.rs index 037e39e5d40..e0c4fd37716 100644 --- a/litellm-rust/crates/cache-redis/src/lib.rs +++ b/litellm-rust/crates/cache-redis/src/lib.rs @@ -4,12 +4,14 @@ pub mod connection; mod counter; mod keys; mod lifecycle; +mod pubsub; mod queue; mod script; mod store; mod topology; pub use cache::RedisCache; +pub use pubsub::RedisSubscription; pub use queue::{RedisLpopOperation, RedisLpopResult, RedisRpushOperation}; pub use script::{RedisArg, RedisScript}; pub use topology::{RedisNode, RedisTopology}; diff --git a/litellm-rust/crates/cache-redis/src/pubsub.rs b/litellm-rust/crates/cache-redis/src/pubsub.rs new file mode 100644 index 00000000000..e64489b93ab --- /dev/null +++ b/litellm-rust/crates/cache-redis/src/pubsub.rs @@ -0,0 +1,154 @@ +use std::{ + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + thread, + time::Duration, +}; + +use litellm_cache::{CacheCodec, Error, Message, MessageStream, PubSubCache}; +use tokio::sync::{mpsc, oneshot}; + +use crate::{cache::RedisCache, topology::RedisTopology}; + +const POLL_INTERVAL: Duration = Duration::from_secs(1); +const QUEUE_CAPACITY: usize = 1024; + +/// Messages arrive on a reader thread that owns the subscribed connection and forwards them +/// into a bounded queue; the thread polls with `POLL_INTERVAL` so `close` and drop can stop it. +pub struct RedisSubscription { + messages: mpsc::Receiver>, + stop: Arc, + reader: Option>, +} + +impl MessageStream for RedisSubscription { + async fn next_message(&mut self, timeout: Option) -> Result, Error> { + let next = match timeout { + Some(limit) if limit.is_zero() => match self.messages.try_recv() { + Ok(message) => Some(message), + Err(mpsc::error::TryRecvError::Empty) => return Ok(None), + Err(mpsc::error::TryRecvError::Disconnected) => None, + }, + Some(limit) => match tokio::time::timeout(limit, self.messages.recv()).await { + Ok(message) => message, + Err(_) => return Ok(None), + }, + None => self.messages.recv().await, + }; + match next { + Some(Ok(message)) => Ok(Some(message)), + Some(Err(error)) => Err(error), + None => Err(Error::Unavailable), + } + } + + async fn close(mut self) -> Result<(), Error> { + self.stop.store(true, Ordering::Release); + let Some(reader) = self.reader.take() else { + return Ok(()); + }; + tokio::task::spawn_blocking(move || reader.join()) + .await + .map_err(|_| Error::Unavailable)? + .map_err(|_| Error::Unavailable) + } +} + +impl Drop for RedisSubscription { + fn drop(&mut self) { + self.stop.store(true, Ordering::Release); + } +} + +impl PubSubCache for RedisCache +where + S: CacheCodec, + C: redis::ConnectionLike + Send + 'static, +{ + type Subscription = RedisSubscription; + + async fn async_publish(&self, channel: &str, payload: &[u8]) -> Result { + if !matches!(self.topology, RedisTopology::Standalone) { + return Err(Error::UnsupportedOperation); + } + let channel = channel.to_owned(); + let payload = payload.to_vec(); + self.run(move |connection| { + redis::cmd("PUBLISH") + .arg(channel) + .arg(payload) + .query::(connection) + .map_err(|_| Error::Unavailable) + }) + .await + } + + async fn async_subscribe(&self, channels: &[String]) -> Result { + let connections = Arc::clone(&self.connections); + let connection = tokio::task::spawn_blocking(move || { + let connection = connections.subscription_connection()?; + connection + .set_read_timeout(Some(POLL_INTERVAL)) + .map_err(|_| Error::Unavailable)?; + Ok::<_, Error>(connection) + }) + .await + .map_err(|_| Error::Unavailable)??; + let (subscribed, ready) = oneshot::channel(); + let (sender, messages) = mpsc::channel(QUEUE_CAPACITY); + let stop = Arc::new(AtomicBool::new(false)); + let reader = thread::Builder::new() + .name("litellm-redis-subscription".into()) + .spawn({ + let stop = Arc::clone(&stop); + let channels = channels.to_vec(); + move || read_messages(connection, &channels, subscribed, &sender, &stop) + }) + .map_err(|_| Error::Unavailable)?; + ready.await.map_err(|_| Error::Unavailable)??; + Ok(RedisSubscription { + messages, + stop, + reader: Some(reader), + }) + } +} + +fn read_messages( + mut connection: redis::Connection, + channels: &[String], + subscribed: oneshot::Sender>, + sender: &mpsc::Sender>, + stop: &AtomicBool, +) { + let mut pubsub = connection.as_pubsub(); + let outcome = channels + .iter() + .try_for_each(|channel| pubsub.subscribe(channel)) + .map_err(|_| Error::Unavailable); + let failed = outcome.is_err(); + let _ = subscribed.send(outcome); + if failed { + return; + } + while !stop.load(Ordering::Acquire) { + match pubsub.get_message() { + Ok(received) => { + let message = Message { + channel: received.get_channel_name().to_owned(), + payload: received.get_payload_bytes().to_vec(), + }; + if sender.blocking_send(Ok(message)).is_err() { + return; + } + } + Err(error) if error.is_timeout() => continue, + Err(_) => { + let _ = sender.blocking_send(Err(Error::Unavailable)); + return; + } + } + } +} diff --git a/litellm-rust/crates/cache-redis/tests/pubsub.rs b/litellm-rust/crates/cache-redis/tests/pubsub.rs new file mode 100644 index 00000000000..9a9d5826a78 --- /dev/null +++ b/litellm-rust/crates/cache-redis/tests/pubsub.rs @@ -0,0 +1,105 @@ +mod support; + +use std::time::Duration; + +use litellm_cache::{Error, JsonCodec, Message, MessageStream, PubSubCache}; +use litellm_cache_redis::RedisCache; +use redis_test::{MockCmd, MockRedisConnection}; +use support::server::PubSubServer; + +type Json = JsonCodec; + +fn live(server: &PubSubServer) -> RedisCache { + RedisCache::new(&server.url(), None, JsonCodec::new()).unwrap() +} + +fn channels(names: &[&str]) -> Vec { + names.iter().map(|name| (*name).to_owned()).collect() +} + +#[tokio::test] +async fn publish_reports_receivers_through_the_pool() { + let connection = MockRedisConnection::new(vec![MockCmd::new( + redis::cmd("PUBLISH") + .arg("events") + .arg(b"payload".as_slice()), + Ok(redis::Value::Int(2)), + )]) + .assert_all_commands_consumed(); + let cache = RedisCache::with_connection(connection, None, Json::new()); + assert_eq!(cache.async_publish("events", b"payload").await, Ok(2)); +} + +#[tokio::test] +async fn subscriptions_need_a_pooled_connection() { + let cache = RedisCache::with_connection(MockRedisConnection::new(vec![]), None, Json::new()); + assert_eq!( + cache + .async_subscribe(&channels(&["events"])) + .await + .map(|_| ()) + .unwrap_err(), + Error::UnsupportedOperation + ); +} + +#[tokio::test] +async fn subscribers_receive_published_messages_until_closed() { + let server = PubSubServer::start(); + let cache = live(&server); + let mut subscription = cache.async_subscribe(&channels(&["events"])).await.unwrap(); + assert_eq!( + subscription.next_message(Some(Duration::ZERO)).await, + Ok(None) + ); + + assert_eq!(cache.async_publish("events", b"hello").await, Ok(1)); + assert_eq!( + subscription + .next_message(Some(Duration::from_secs(5))) + .await, + Ok(Some(Message { + channel: "events".into(), + payload: b"hello".to_vec(), + })) + ); + assert_eq!(cache.async_publish("other", b"ignored").await, Ok(0)); + assert_eq!( + subscription + .next_message(Some(Duration::from_millis(50))) + .await, + Ok(None) + ); + + subscription.close().await.unwrap(); + assert_eq!(server.unsubscribed(), vec!["events".to_owned()]); + assert_eq!(cache.async_publish("events", b"nobody").await, Ok(0)); +} + +#[tokio::test] +async fn a_dropped_connection_surfaces_as_unavailable() { + let server = PubSubServer::start(); + let cache = live(&server); + let mut subscription = cache.async_subscribe(&channels(&["events"])).await.unwrap(); + server.drop_subscribers(); + assert_eq!( + subscription + .next_message(Some(Duration::from_secs(5))) + .await, + Err(Error::Unavailable) + ); +} + +#[tokio::test] +async fn subscribing_to_an_unreachable_server_fails_before_returning_a_stream() { + let url = PubSubServer::start().url(); + let cache: RedisCache = RedisCache::new(&url, None, JsonCodec::new()).unwrap(); + assert_eq!( + cache + .async_subscribe(&channels(&["events"])) + .await + .map(|_| ()) + .unwrap_err(), + Error::Unavailable + ); +} diff --git a/litellm-rust/crates/cache-redis/tests/support/mod.rs b/litellm-rust/crates/cache-redis/tests/support/mod.rs index bff1e618582..f2b3b5e4061 100644 --- a/litellm-rust/crates/cache-redis/tests/support/mod.rs +++ b/litellm-rust/crates/cache-redis/tests/support/mod.rs @@ -1,5 +1,7 @@ #![allow(dead_code)] +pub mod server; + use std::{ collections::BTreeMap, time::{Duration, SystemTime, UNIX_EPOCH}, diff --git a/litellm-rust/crates/cache-redis/tests/support/server.rs b/litellm-rust/crates/cache-redis/tests/support/server.rs new file mode 100644 index 00000000000..7f362e11f68 --- /dev/null +++ b/litellm-rust/crates/cache-redis/tests/support/server.rs @@ -0,0 +1,253 @@ +use std::{ + collections::HashMap, + io::{Read, Write}, + net::{Shutdown, SocketAddr, TcpListener, TcpStream}, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, AtomicU64, Ordering}, + }, + thread, +}; + +#[derive(Default)] +struct Registry { + subscribers: HashMap>, + unsubscribed: Vec, +} + +/// A RESP2 server over a real socket that speaks PING, SUBSCRIBE, UNSUBSCRIBE and PUBLISH, so +/// the subscription path runs against the same wire protocol as a live Redis. +pub struct PubSubServer { + address: SocketAddr, + registry: Arc>, + shutdown: Arc, + acceptor: Option>, +} + +impl PubSubServer { + pub fn start() -> Self { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + let registry = Arc::new(Mutex::new(Registry::default())); + let shutdown = Arc::new(AtomicBool::new(false)); + let acceptor = thread::spawn({ + let registry = Arc::clone(®istry); + let shutdown = Arc::clone(&shutdown); + let next_id = AtomicU64::new(1); + move || { + for stream in listener.incoming() { + if shutdown.load(Ordering::Acquire) { + break; + } + let Ok(stream) = stream else { break }; + let id = next_id.fetch_add(1, Ordering::Relaxed); + let registry = Arc::clone(®istry); + thread::spawn(move || serve(stream, id, ®istry)); + } + } + }); + Self { + address, + registry, + shutdown, + acceptor: Some(acceptor), + } + } + + pub fn url(&self) -> String { + format!("redis://{}", self.address) + } + + pub fn unsubscribed(&self) -> Vec { + self.registry.lock().unwrap().unsubscribed.clone() + } + + /// Closes every subscriber socket, as a restarted server would. + pub fn drop_subscribers(&self) { + let mut registry = self.registry.lock().unwrap(); + for (_, stream) in registry + .subscribers + .drain() + .flat_map(|(_, subscribers)| subscribers) + { + let _ = stream.shutdown(Shutdown::Both); + } + } +} + +impl Drop for PubSubServer { + fn drop(&mut self) { + self.shutdown.store(true, Ordering::Release); + let _ = TcpStream::connect(self.address); + if let Some(acceptor) = self.acceptor.take() { + let _ = acceptor.join(); + } + } +} + +fn serve(mut stream: TcpStream, id: u64, registry: &Mutex) { + let mut buffer = Vec::new(); + let mut chunk = [0u8; 4096]; + loop { + let read = match stream.read(&mut chunk) { + Ok(0) | Err(_) => break, + Ok(read) => read, + }; + buffer.extend_from_slice(&chunk[..read]); + while let Some((command, consumed)) = parse_command(&buffer) { + buffer.drain(..consumed); + let reply = respond(&command, &stream, id, registry); + if stream.write_all(&reply).is_err() { + return; + } + } + } + let mut registry = registry.lock().unwrap(); + for subscribers in registry.subscribers.values_mut() { + subscribers.retain(|(subscriber, _)| *subscriber != id); + } +} + +fn respond( + command: &[Vec], + stream: &TcpStream, + id: u64, + registry: &Mutex, +) -> Vec { + let name = String::from_utf8_lossy(&command[0]).to_ascii_uppercase(); + let text = |bytes: &[u8]| String::from_utf8_lossy(bytes).into_owned(); + match name.as_str() { + "PING" => b"+PONG\r\n".to_vec(), + "SUBSCRIBE" => { + let mut registry = registry.lock().unwrap(); + command[1..] + .iter() + .flat_map(|channel| { + let channel = text(channel); + registry + .subscribers + .entry(channel.clone()) + .or_default() + .push((id, stream.try_clone().unwrap())); + let count = subscriptions(®istry, id); + array(&[bulk(b"subscribe"), bulk(channel.as_bytes()), integer(count)]) + }) + .collect() + } + "UNSUBSCRIBE" | "PUNSUBSCRIBE" => { + let mut registry = registry.lock().unwrap(); + let named: Vec = command[1..].iter().map(|channel| text(channel)).collect(); + let channels: Vec = if named.is_empty() { + registry + .subscribers + .iter() + .filter(|(_, subscribers)| subscribers.iter().any(|(sub, _)| *sub == id)) + .map(|(channel, _)| channel.clone()) + .collect() + } else { + named + }; + if channels.is_empty() { + return array(&[ + bulk(&name.to_ascii_lowercase().into_bytes()), + b"$-1\r\n".to_vec(), + integer(0), + ]); + } + channels + .into_iter() + .flat_map(|channel| { + if let Some(subscribers) = registry.subscribers.get_mut(&channel) { + let before = subscribers.len(); + subscribers.retain(|(sub, _)| *sub != id); + if subscribers.len() < before { + registry.unsubscribed.push(channel.clone()); + } + } + let count = subscriptions(®istry, id); + array(&[ + bulk(b"unsubscribe"), + bulk(channel.as_bytes()), + integer(count), + ]) + }) + .collect() + } + "PUBLISH" => { + let channel = text(&command[1]); + let payload = &command[2]; + let mut registry = registry.lock().unwrap(); + let delivered = registry + .subscribers + .get_mut(&channel) + .map_or(0, |subscribers| { + subscribers.retain_mut(|(_, subscriber)| { + subscriber + .write_all(&array(&[ + bulk(b"message"), + bulk(channel.as_bytes()), + bulk(payload), + ])) + .is_ok() + }); + subscribers.len() + }); + integer(delivered) + } + _ => b"+OK\r\n".to_vec(), + } +} + +fn subscriptions(registry: &Registry, id: u64) -> usize { + registry + .subscribers + .values() + .filter(|subscribers| subscribers.iter().any(|(sub, _)| *sub == id)) + .count() +} + +fn parse_command(buffer: &[u8]) -> Option<(Vec>, usize)> { + fn line(buffer: &[u8], start: usize) -> Option<(&[u8], usize)> { + let end = buffer[start..] + .windows(2) + .position(|window| window == b"\r\n")? + + start; + Some((&buffer[start..end], end + 2)) + } + fn length(line: &[u8]) -> Option { + std::str::from_utf8(line.get(1..)?).ok()?.parse().ok() + } + let (header, mut cursor) = line(buffer, 0)?; + let count = length(header)?; + let mut command = Vec::with_capacity(count); + for _ in 0..count { + let (size, start) = line(buffer, cursor)?; + let size = length(size)?; + let argument = buffer.get(start..start + size)?; + command.push(argument.to_vec()); + cursor = start + size + 2; + if buffer.len() < cursor { + return None; + } + } + Some((command, cursor)) +} + +fn bulk(bytes: &[u8]) -> Vec { + let mut out = format!("${}\r\n", bytes.len()).into_bytes(); + out.extend_from_slice(bytes); + out.extend_from_slice(b"\r\n"); + out +} + +fn integer(value: usize) -> Vec { + format!(":{value}\r\n").into_bytes() +} + +fn array(items: &[Vec]) -> Vec { + let mut out = format!("*{}\r\n", items.len()).into_bytes(); + for item in items { + out.extend_from_slice(item); + } + out +} diff --git a/litellm-rust/crates/cache/src/capabilities.rs b/litellm-rust/crates/cache/src/capabilities.rs index ac9d8fc8764..e19917ffff3 100644 --- a/litellm-rust/crates/cache/src/capabilities.rs +++ b/litellm-rust/crates/cache/src/capabilities.rs @@ -303,3 +303,38 @@ pub trait ScriptCache: BaseCache { fn async_register_script(&self, source: String) -> Self::Script; } + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct Message { + pub channel: String, + pub payload: Vec, +} + +/// One open subscription. `next_message` returns `None` when `timeout` elapses first, and +/// `Err` once the connection behind it is gone, so the owner reconnects with a fresh +/// `async_subscribe`. +pub trait MessageStream: Send { + fn next_message( + &mut self, + timeout: Option, + ) -> impl Future, Error>> + Send; + + fn close(self) -> impl Future> + Send; +} + +/// `async_publish` and `async_subscribe`. Channels are passed through as given; a namespace +/// prefix is the caller's, as with Python's `f"{namespace}:{channel}"`. +pub trait PubSubCache: BaseCache { + type Subscription: MessageStream; + + fn async_publish( + &self, + channel: &str, + payload: &[u8], + ) -> impl Future> + Send; + + fn async_subscribe( + &self, + channels: &[String], + ) -> impl Future> + Send; +} diff --git a/litellm-rust/crates/cache/src/lib.rs b/litellm-rust/crates/cache/src/lib.rs index 55d8a5fa28b..9f1a74874a9 100644 --- a/litellm-rust/crates/cache/src/lib.rs +++ b/litellm-rust/crates/cache/src/lib.rs @@ -16,8 +16,9 @@ pub use caching::{Cache, CacheBackend, get_cache, set_cache}; pub use capabilities::{ BatchCache, BoundedCounterCache, BulkDeleteCache, CacheScript, ClaimCache, ClientInfoCache, ConnectionCache, CountReadCache, CounterCache, DeleteCache, DisconnectCache, FlushAllCache, - FlushCache, IncrementOperation, PingCache, PopOperation, PushOperation, QueueCache, - RefreshTtlCache, ScanCache, ScriptCache, SetCache, TtlCache, TtlPipelineCache, + FlushCache, IncrementOperation, Message, MessageStream, PingCache, PopOperation, PubSubCache, + PushOperation, QueueCache, RefreshTtlCache, ScanCache, ScriptCache, SetCache, TtlCache, + TtlPipelineCache, }; pub use codec::{CacheCodec, JsonCodec}; pub use dual::{ diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 7cec84e0ebb..690d37551fb 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -17,13 +17,14 @@ import json import logging import threading import time -from collections.abc import Awaitable, Callable, Iterator, Sequence +from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence from contextvars import ContextVar from dataclasses import dataclass from datetime import timedelta from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast -from pydantic import TypeAdapter +from pydantic import TypeAdapter, ValidationError +from typing_extensions import ReadOnly, TypedDict import litellm from litellm._logging import print_verbose, verbose_logger @@ -52,6 +53,7 @@ if TYPE_CHECKING: from prometheus_client import Gauge as _PromGauge from redis.asyncio import Redis, RedisCluster from redis.asyncio.client import Pipeline + from redis.asyncio.client import PubSub as AsyncPubSub from redis.asyncio.cluster import ClusterPipeline pipeline = Pipeline @@ -577,6 +579,63 @@ def _redis_circuit_breaker_guard_sync(method: Callable[..., _RedisCallResult]) - ) +class _PubSubFrame(TypedDict): + type: ReadOnly[str] + channel: ReadOnly[str | bytes] + data: ReadOnly[str | bytes | int] + + +_PUBSUB_FRAME: Final = TypeAdapter(_PubSubFrame) +_PUBSUB_WAIT_SLICE_SECONDS: Final = 1.0 + + +@dataclass(frozen=True, slots=True) +class RedisMessage: + channel: str + payload: bytes + + +def _redis_message(frame: object) -> RedisMessage | None: + """The published message a pub/sub frame carries; ``None`` for subscribe acks and anything malformed.""" + try: + parsed: Final = _PUBSUB_FRAME.validate_python(frame) + except ValidationError: + return None + channel: Final = parsed["channel"] + payload: Final = parsed["data"] + if parsed["type"] not in ("message", "pmessage") or isinstance(payload, int): + return None + return RedisMessage( + channel=channel.decode("utf-8", errors="replace") if isinstance(channel, bytes) else channel, + payload=payload.encode("utf-8") if isinstance(payload, str) else payload, + ) + + +@dataclass(frozen=True, slots=True) +class RedisSubscription: + """One SUBSCRIBE session. ``get_message`` waits up to ``timeout`` seconds (``None`` waits + indefinitely) for a published message and skips the acks redis-py interleaves with them.""" + + pubsub: "AsyncPubSub" + + async def get_message(self, *, timeout: float | None) -> RedisMessage | None: + clock: Final = asyncio.get_running_loop().time + deadline: Final = None if timeout is None else clock() + timeout + while True: + remaining = ( # rebind-ok: countdown per frame + _PUBSUB_WAIT_SLICE_SECONDS if deadline is None else max(deadline - clock(), 0.0) + ) + frame: object = await self.pubsub.get_message(timeout=remaining) # rebind-ok: one frame per loop turn + if frame is None: + return None + message = _redis_message(frame) # rebind-ok: one frame per loop turn + if message is not None: + return message + + async def aclose(self) -> None: + await self.pubsub.aclose() # pyright: ignore[reportAttributeAccessIssue] # types-redis 4.6 stubs predate PubSub.aclose + + class RedisCache(BaseCache): # if users don't provider one, use the default litellm cache @@ -1839,6 +1898,33 @@ class RedisCache(BaseCache): key = self.check_and_fix_namespace(key=key) self.redis_client.delete(key) + def _standalone_async_client(self) -> "Redis[bytes] | Redis[str]": + from redis.asyncio import RedisCluster + + client: Final[object] = self.init_async_client() + if isinstance(client, RedisCluster): + raise NotImplementedError("Redis Cluster clients have no pub/sub support") + return cast("Redis[bytes] | Redis[str]", client) # cast-ok: decode_responses picks the reply type at runtime + + @_redis_circuit_breaker_guard + async def async_publish(self, channel: str, message: str | bytes) -> int: + return await self._standalone_async_client().publish(channel, message) + + @_redis_circuit_breaker_guard + async def async_subscribe(self, *channels: str) -> RedisSubscription: + pubsub: Final = self._standalone_async_client().pubsub() + await pubsub.subscribe(*channels) + return RedisSubscription(pubsub) + + def connection_pool_status(self) -> Mapping[str, object]: + pool: Final = getattr(self.redis_client, "connection_pool", None) + if pool is None: + return {} + return { + "max_connections": getattr(pool, "max_connections", None), + "connection_class": getattr(getattr(pool, "connection_class", None), "__name__", None), + } + async def _pipeline_increment_helper( self, pipe: pipeline, diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_distributed_lock.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_distributed_lock.py index c7eefed75ed..615337d3831 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_distributed_lock.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/redis_distributed_lock.py @@ -1,94 +1,84 @@ -"""Concrete ``DistributedLock`` over a Redis client: ``SET NX PX`` / owner-only renew / delete. +"""Concrete ``DistributedLock`` over ``RedisCache`` scripts: ``SET NX PX`` / owner-only renew / delete. The cross-replica lock the ``RedisRefreshCoordinator`` elects refreshers with. ``acquire`` is an atomic ``SET key token NX PX ttl`` (only the first caller wins; the entry self-expires so a crashed holder can't wedge refresh). ``extend`` renews the lease only when the token still matches, and ``release`` deletes the key only when it still holds this caller's token, so a holder whose lock already PX-expired and was re-acquired by another worker cannot delete the new holder's lock. ``is_held`` is -``EXISTS``. Every key is run through the injected ``namespace_key`` before it reaches Redis, so lock -keys carry the same namespace as cache keys and cannot collide with another deployment sharing Redis. +``EXISTS``. Every operation runs through ``RedisCache.async_register_script``, which prefixes the key +with the cache namespace, so lock keys carry the same namespace as cache keys and cannot collide with +another deployment sharing Redis. -The Redis client is injected (in production the async client from LiteLLM's ``RedisCache``), so the -lock is unit-testable with a fake. A transport error on ``acquire`` returns ``LockAcquisition.ERROR`` - -distinct from ``HELD`` - so the coordinator refreshes anyway instead of mistaking a dead backend for a -busy holder; a Redis blip degrades to an extra refresh, never a stale bearer. +A transport error on ``acquire`` returns ``LockAcquisition.ERROR`` - distinct from ``HELD`` - so the +coordinator refreshes anyway instead of mistaking a dead backend for a busy holder; a Redis blip +degrades to an extra refresh, never a stale bearer. """ from __future__ import annotations -from collections.abc import Callable -from dataclasses import KW_ONLY, dataclass -from typing import Final, Protocol +from dataclasses import dataclass +from typing import TYPE_CHECKING, Final + +from pydantic import TypeAdapter from litellm._logging import verbose_logger from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import ( LockAcquisition, ) -# Delete the key only if it still holds this caller's token, so a holder whose lock already expired -# (PX) and was re-acquired by another worker cannot delete the new holder's lock. -_RELEASE_IF_OWNER = "if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) else return 0 end" +if TYPE_CHECKING: + from litellm.caching.redis_cache import RedisCache + +_ACQUIRE_IF_ABSENT: Final = "return redis.call('set', KEYS[1], ARGV[1], 'NX', 'PX', ARGV[2])" _EXTEND_IF_OWNER: Final = ( "if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('pexpire', KEYS[1], ARGV[2]) else return 0 end" ) +_RELEASE_IF_OWNER: Final = ( + "if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) else return 0 end" +) +_IS_HELD: Final = "return redis.call('exists', KEYS[1])" +_SCRIPT_REPLY: Final[TypeAdapter[int | bytes | str | None]] = TypeAdapter(int | bytes | str | None) -class RedisCommands(Protocol): - """The slice of the async Redis client this lock needs.""" - - async def set(self, name: str, value: str, *, nx: bool = False, px: int | None = None) -> object | None: ... - - async def eval(self, script: str, numkeys: int, *keys_and_args: str) -> object: ... - - async def exists(self, *names: str) -> int: ... +def _milliseconds(seconds: float) -> str: + return str(int(seconds * 1000)) @dataclass(frozen=True, slots=True) class RedisDistributedLock: - client: RedisCommands - _: KW_ONLY - namespace_key: Callable[[str], str] = lambda key: key + redis_cache: RedisCache + + async def _run(self, script: str, key: str, *args: str) -> int | bytes | str | None: + return _SCRIPT_REPLY.validate_python( + await self.redis_cache.async_register_script(script)(keys=(key,), args=args) + ) async def acquire(self, key: str, token: str, ttl_seconds: float) -> LockAcquisition: try: - result: Final = await self.client.set(self.namespace_key(key), token, nx=True, px=int(ttl_seconds * 1000)) - # Degrade on any Redis client error: redis.exceptions narrows only via an import that - # is Unknown under basedpyright, and the lock must never crash the resolve path. - except Exception as exc: # noqa: BLE001 + result: Final = await self._run(_ACQUIRE_IF_ABSENT, key, token, _milliseconds(ttl_seconds)) + except Exception as exc: # noqa: BLE001 # the lock must never crash the resolve path verbose_logger.warning("RedisDistributedLock.acquire failed: %s", exc) return LockAcquisition.ERROR return LockAcquisition.ACQUIRED if result is not None else LockAcquisition.HELD async def extend(self, key: str, token: str, ttl_seconds: float) -> bool: try: - result: Final = await self.client.eval( - _EXTEND_IF_OWNER, - 1, - self.namespace_key(key), - token, - str(int(ttl_seconds * 1000)), - ) - # Degrade on any Redis client error: redis.exceptions narrows only via an import that - # is Unknown under basedpyright, and the lock must never crash the resolve path. - except Exception as exc: # noqa: BLE001 + result: Final = await self._run(_EXTEND_IF_OWNER, key, token, _milliseconds(ttl_seconds)) + except Exception as exc: # noqa: BLE001 # the lock must never crash the resolve path verbose_logger.warning("RedisDistributedLock.extend failed: %s", exc) return False return result == 1 async def release(self, key: str, token: str) -> None: try: - await self.client.eval(_RELEASE_IF_OWNER, 1, self.namespace_key(key), token) - # Degrade on any Redis client error: redis.exceptions narrows only via an import that - # is Unknown under basedpyright, and the lock must never crash the resolve path. - except Exception as exc: # noqa: BLE001 + await self._run(_RELEASE_IF_OWNER, key, token) + except Exception as exc: # noqa: BLE001 # the lock must never crash the resolve path verbose_logger.warning("RedisDistributedLock.release failed: %s", exc) async def is_held(self, key: str) -> bool: try: - return await self.client.exists(self.namespace_key(key)) > 0 - # Degrade on any Redis client error: redis.exceptions narrows only via an import that - # is Unknown under basedpyright, and the lock must never crash the resolve path. - except Exception as exc: # noqa: BLE001 - # On error, report "not held" so a waiter stops waiting and re-reads rather than blocking. + result: Final = await self._run(_IS_HELD, key) + except Exception as exc: # noqa: BLE001 # a waiter stops waiting and re-reads rather than blocking verbose_logger.warning("RedisDistributedLock.is_held failed: %s", exc) return False + return isinstance(result, int) and result > 0 diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/runtime_refresh_coordinator.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/runtime_refresh_coordinator.py index e799838b5b7..1c2a7230a44 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/runtime_refresh_coordinator.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/runtime_refresh_coordinator.py @@ -31,11 +31,4 @@ def runtime_refresh_coordinator() -> RefreshCoordinator | None: redis_cache: Final = user_api_key_cache.redis_cache if redis_cache is None: return None - # The Redis client from init_async_client() is only partially typed; the lock validates every - # reply it depends on, so the untyped boundary is contained here. - redis_client: Final = redis_cache.init_async_client() # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType] # litellm redis wrapper is untyped - lock: Final = RedisDistributedLock( - redis_client, # pyright: ignore[reportArgumentType,reportUnknownArgumentType] # litellm redis wrapper is untyped - namespace_key=redis_cache.check_and_fix_namespace, - ) - return RedisRefreshCoordinator(lock) + return RedisRefreshCoordinator(RedisDistributedLock(redis_cache)) diff --git a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py index 11cb66d1a7f..bc793c94187 100644 --- a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py +++ b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py @@ -5,15 +5,11 @@ from dataclasses import asdict, dataclass from typing import TYPE_CHECKING, Final from litellm._logging import verbose_proxy_logger -from litellm.proxy.common_utils.config_sync_pubsub import ( - _ConfigSyncPubSub, - _pubsub_capable_client, - coordination_redis_cache, -) +from litellm.proxy.common_utils.config_sync_pubsub import coordination_redis_cache if TYPE_CHECKING: from litellm.caching.in_memory_cache import InMemoryCache - from litellm.caching.redis_cache import RedisCache + from litellm.caching.redis_cache import RedisCache, RedisMessage, RedisSubscription from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache AUTH_CACHE_INVALIDATION_CHANNEL: Final = "litellm_proxy.auth_cache_invalidation" @@ -73,15 +69,13 @@ def _message_from_data(data: object) -> _CacheInvalidationMessage | None: async def _publish_to_redis(redis_cache: "RedisCache", cache_key: str, message: str) -> None: try: - client: Final = _pubsub_capable_client(redis_cache) - if client is None: - verbose_proxy_logger.debug( - "auth cache invalidation publish for %s skipped: cluster redis client has no pub/sub support", - cache_key, - ) - return async with _in_flight_publishes: - await client.publish(auth_cache_invalidation_channel(redis_cache), message) + await redis_cache.async_publish(auth_cache_invalidation_channel(redis_cache), message) + except NotImplementedError: + verbose_proxy_logger.debug( + "auth cache invalidation publish for %s skipped: cluster redis client has no pub/sub support", + cache_key, + ) except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors verbose_proxy_logger.warning("auth cache invalidation publish for %s failed: %s", cache_key, e) @@ -184,20 +178,14 @@ class AuthCacheInvalidationSubscriber: backoff_seconds = _BACKOFF_INITIAL_SECONDS # rebind-ok: exponential backoff accumulator across reconnects while True: try: - client = _pubsub_capable_client(self._redis_cache) - if client is None: - verbose_proxy_logger.warning( - "auth cache invalidation subscriber disabled: cluster redis client has no pub/sub support; " - "cross-worker eviction falls back to the local cache TTL" - ) + subscription = await self._open_subscription() + if subscription is None: return - pubsub = client.pubsub() try: - await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache)) backoff_seconds = _BACKOFF_INITIAL_SECONDS - await self._consume(pubsub) + await self._consume(subscription) finally: - await self._close_pubsub(pubsub) + await self._close_subscription(subscription) except asyncio.CancelledError: raise except Exception as e: # noqa: BLE001 # any redis failure falls through to backoff and reconnect @@ -209,16 +197,25 @@ class AuthCacheInvalidationSubscriber: await asyncio.sleep(backoff_seconds) backoff_seconds = min(backoff_seconds * 2, _BACKOFF_MAX_SECONDS) - async def _consume(self, pubsub: _ConfigSyncPubSub) -> None: + async def _open_subscription(self) -> "RedisSubscription | None": + try: + return await self._redis_cache.async_subscribe(auth_cache_invalidation_channel(self._redis_cache)) + except NotImplementedError: + verbose_proxy_logger.warning( + "auth cache invalidation subscriber disabled: cluster redis client has no pub/sub support; " + "cross-worker eviction falls back to the local cache TTL" + ) + return None + + async def _consume(self, subscription: "RedisSubscription") -> None: while True: - message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS) + message = await subscription.get_message(timeout=_POLL_TIMEOUT_SECONDS) if message is None: continue self._apply_message(message) - def _apply_message(self, message: object) -> None: - data: Final = message.get("data") if isinstance(message, dict) else None - parsed: Final = _message_from_data(data) + def _apply_message(self, message: "RedisMessage") -> None: + parsed: Final = _message_from_data(message.payload) if parsed is None: return if parsed.new_value is not None: @@ -230,8 +227,8 @@ class AuthCacheInvalidationSubscriber: additional_cache.delete_cache(parsed.cache_key) @staticmethod - async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None: + async def _close_subscription(subscription: "RedisSubscription") -> None: try: - await pubsub.aclose() + await subscription.aclose() except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection verbose_proxy_logger.debug("auth cache invalidation pubsub close failed: %s", e) diff --git a/litellm/proxy/common_utils/config_sync_pubsub.py b/litellm/proxy/common_utils/config_sync_pubsub.py index b20c0d9c9a5..094ff6d3bfb 100644 --- a/litellm/proxy/common_utils/config_sync_pubsub.py +++ b/litellm/proxy/common_utils/config_sync_pubsub.py @@ -4,27 +4,13 @@ import random import time from collections.abc import Awaitable, Callable from dataclasses import asdict, dataclass -from typing import TYPE_CHECKING, Final, Protocol, cast # noqa: TID251 # untyped prisma/redis boundary needs cast +from typing import TYPE_CHECKING, Final, cast # noqa: TID251 # untyped prisma boundary needs cast from litellm._logging import verbose_proxy_logger from litellm.repositories.prisma_protocols import RowT_co, TableActions if TYPE_CHECKING: - from litellm.caching.redis_cache import RedisCache - - -class _ConfigSyncPubSub(Protocol): - def subscribe(self, *channels: str) -> Awaitable[object]: ... - - def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Awaitable[object]: ... - - def aclose(self) -> Awaitable[object]: ... - - -class _ConfigSyncPubSubClient(Protocol): - def publish(self, channel: str, message: str) -> Awaitable[int]: ... - - def pubsub(self) -> _ConfigSyncPubSub: ... + from litellm.caching.redis_cache import RedisCache, RedisSubscription CONFIG_SYNC_CHANNEL: Final = "litellm_proxy.config_change" @@ -82,22 +68,6 @@ def config_sync_channel(redis_cache: "RedisCache") -> str: return f"{redis_cache.namespace}:{CONFIG_SYNC_CHANNEL}" -def _raw_async_client(redis_cache: "RedisCache") -> object: - return cast( # cast-ok: redis-py generics leave the client type partially unknown - object, - redis_cache.init_async_client(), # pyright: ignore[reportUnknownMemberType] # redis generics - ) - - -def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient | None: - from redis.asyncio import Redis - - client: Final = _raw_async_client(redis_cache) - if isinstance(client, Redis): - return cast(_ConfigSyncPubSubClient, client) # cast-ok: protocol view of the standalone redis client - return None - - @dataclass(frozen=True, slots=True) class _ConfigChangeMessage: object_type: str @@ -111,14 +81,12 @@ async def publish_config_change(redis_cache: "RedisCache | None", object_type: s if redis_cache is None: return try: - client: Final = _pubsub_capable_client(redis_cache) - if client is None: - verbose_proxy_logger.debug( - "config sync publish for %s skipped: cluster redis client has no pub/sub support", - object_type, - ) - return - await client.publish(config_sync_channel(redis_cache), _config_change_message_json(object_type)) + await redis_cache.async_publish(config_sync_channel(redis_cache), _config_change_message_json(object_type)) + except NotImplementedError: + verbose_proxy_logger.debug( + "config sync publish for %s skipped: cluster redis client has no pub/sub support", + object_type, + ) except Exception as e: # noqa: BLE001 # best-effort publish; writes must never fail on redis errors verbose_proxy_logger.warning("config sync publish for %s failed: %s", object_type, e) @@ -237,20 +205,14 @@ class ConfigSyncSubscriber: backoff_seconds = self._backoff_initial_seconds while True: try: - client = _pubsub_capable_client(self._redis_cache) - if client is None: - verbose_proxy_logger.warning( - "config sync subscriber disabled: cluster redis client has no pub/sub support; " - "interval polling remains the only sync mechanism" - ) + subscription = await self._open_subscription() + if subscription is None: return - pubsub = client.pubsub() try: - await pubsub.subscribe(config_sync_channel(self._redis_cache)) backoff_seconds = self._backoff_initial_seconds - await self._consume(pubsub) + await self._consume(subscription) finally: - await self._close_pubsub(pubsub) + await self._close_subscription(subscription) except asyncio.CancelledError: raise except Exception as e: # noqa: BLE001 # any redis failure falls through to backoff and reconnect @@ -262,14 +224,24 @@ class ConfigSyncSubscriber: await self._sleep(backoff_seconds) backoff_seconds = min(backoff_seconds * 2, self._backoff_max_seconds) - async def _consume(self, pubsub: _ConfigSyncPubSub) -> None: + async def _open_subscription(self) -> "RedisSubscription | None": + try: + return await self._redis_cache.async_subscribe(config_sync_channel(self._redis_cache)) + except NotImplementedError: + verbose_proxy_logger.warning( + "config sync subscriber disabled: cluster redis client has no pub/sub support; " + "interval polling remains the only sync mechanism" + ) + return None + + async def _consume(self, subscription: "RedisSubscription") -> None: while True: - message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS) + message = await subscription.get_message(timeout=_POLL_TIMEOUT_SECONDS) if message is None: continue await self._sleep(self._debounce_seconds + self._rng.uniform(0.0, self._jitter_max_seconds)) await self._wait_for_min_resync_interval() - await self._drain_pending(pubsub) + await self._drain_pending(subscription) await self._run_resync_callbacks() self._last_resync_at = self._monotonic() @@ -286,8 +258,8 @@ class ConfigSyncSubscriber: await self._sleep(seconds_until_next_resync) @staticmethod - async def _drain_pending(pubsub: _ConfigSyncPubSub) -> None: - while await pubsub.get_message(ignore_subscribe_messages=True, timeout=0) is not None: + async def _drain_pending(subscription: "RedisSubscription") -> None: + while await subscription.get_message(timeout=0) is not None: pass async def _run_resync_callbacks(self) -> None: @@ -298,8 +270,8 @@ class ConfigSyncSubscriber: verbose_proxy_logger.warning("config sync resync callback failed: %s", e) @staticmethod - async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None: + async def _close_subscription(subscription: "RedisSubscription") -> None: try: - await pubsub.aclose() + await subscription.aclose() except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection verbose_proxy_logger.debug("config sync pubsub close failed: %s", e) diff --git a/litellm/proxy/common_utils/debug_utils.py b/litellm/proxy/common_utils/debug_utils.py index 2544321a1b6..7324b8e7de1 100644 --- a/litellm/proxy/common_utils/debug_utils.py +++ b/litellm/proxy/common_utils/debug_utils.py @@ -8,7 +8,7 @@ import sys import tracemalloc from collections import Counter from collections.abc import Mapping, Sequence -from typing import Annotated, Any, Final, NamedTuple, Protocol, TypedDict +from typing import TYPE_CHECKING, Annotated, Any, Final, NamedTuple, Protocol, TypedDict from fastapi import APIRouter, Depends, HTTPException, Query from typing_extensions import ReadOnly @@ -22,6 +22,9 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.bug_report_config import build_proxy_environment_report from litellm.proxy.common_utils.resource_ownership import is_proxy_admin +if TYPE_CHECKING: + from litellm.caching.redis_cache import RedisCache + router: Final = APIRouter() @@ -482,7 +485,7 @@ def _get_uncollectable_objects_info() -> Mapping[str, object]: def _get_cache_memory_stats( - user_api_key_cache, llm_router, proxy_logging_obj, redis_usage_cache + user_api_key_cache, llm_router, proxy_logging_obj, redis_usage_cache: "RedisCache | None" ) -> Mapping[str, object]: """Calculate memory usage for all caches.""" cache_stats: Final[dict[str, object]] = {} @@ -529,22 +532,8 @@ def _get_cache_memory_stats( cache_stats["redis_usage_cache"] = { "enabled": True, "cache_type": type(redis_usage_cache).__name__, + "connection_pool": redis_usage_cache.connection_pool_status(), } - # Try to get Redis connection pool info if available - try: - if hasattr(redis_usage_cache, "redis_client") and redis_usage_cache.redis_client: - if hasattr(redis_usage_cache.redis_client, "connection_pool"): - pool_info: Final = redis_usage_cache.redis_client.connection_pool - cache_stats["redis_usage_cache"]["connection_pool"] = { - "max_connections": ( - pool_info.max_connections if hasattr(pool_info, "max_connections") else None - ), - "connection_class": ( - pool_info.connection_class.__name__ if hasattr(pool_info, "connection_class") else None - ), - } - except Exception as e: - verbose_proxy_logger.debug("Error getting Redis pool info: %s", e) else: cache_stats["redis_usage_cache"] = {"enabled": False} diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index 5d72fe7213d..73b7004d18b 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -58,9 +58,7 @@ def test_check_and_fix_namespace_prefixes_keys_sharing_the_namespace_prefix( @pytest.mark.parametrize("namespace", [None, "litellm"]) @pytest.mark.asyncio -async def test_async_delete_cache_applies_namespace( - namespace, monkeypatch, redis_no_ping -): +async def test_async_delete_cache_applies_namespace(namespace, monkeypatch, redis_no_ping): """async_delete_cache must prefix keys with the namespace, matching every other cache operation. Without this, Redis NOPERM errors occur when an ACL restricts DEL to the litellm:* pattern.""" @@ -68,9 +66,7 @@ async def test_async_delete_cache_applies_namespace( redis_cache = RedisCache(namespace=namespace) mock_redis_instance = AsyncMock() - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): await redis_cache.async_delete_cache(key="3997c4abcdef") expected_key = "litellm:3997c4abcdef" if namespace else "3997c4abcdef" @@ -133,9 +129,7 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch): ] # Test the helper method - result = await redis_cache.handle_lpop_count_for_older_redis_versions( - pipe=mock_pipeline, key="test_key", count=2 - ) + result = await redis_cache.handle_lpop_count_for_older_redis_versions(pipe=mock_pipeline, key="test_key", count=2) # Verify results assert result == [b"value1", b"value2"] @@ -144,18 +138,14 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch): @pytest.mark.asyncio -async def test_async_rpush_pipeline_empty_list_returns_empty( - monkeypatch, redis_no_ping -): +async def test_async_rpush_pipeline_empty_list_returns_empty(monkeypatch, redis_no_ping): """Empty rpush_list should return empty list without touching Redis""" monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache() mock_redis_instance = AsyncMock() - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): result = await redis_cache.async_rpush_pipeline(rpush_list=[]) assert result == [] @@ -170,9 +160,7 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping): mock_redis_instance = AsyncMock() - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): result = await redis_cache.async_lpop_pipeline(lpop_list=[]) assert result == [] @@ -197,9 +185,7 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping): ], ) @pytest.mark.asyncio -async def test_async_register_script_namespaces_keys( - namespace, raw_keys, expected_keys, monkeypatch, redis_no_ping -): +async def test_async_register_script_namespaces_keys(namespace, raw_keys, expected_keys, monkeypatch, redis_no_ping): """The callable returned by async_register_script (used by the rate limiter Lua scripts, pod-lock release, and budget limiters) must namespace every key it is invoked with. The hash tag is preserved so cluster slotting is intact.""" @@ -210,16 +196,12 @@ async def test_async_register_script_namespaces_keys( mock_redis_instance = MagicMock() mock_redis_instance.register_script = MagicMock(return_value=registered_script) - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): script = redis_cache.async_register_script("return 1") result = await script(keys=raw_keys, args=[60]) assert result == "ok" - registered_script.assert_awaited_once_with( - keys=tuple(expected_keys), args=[60], client=None - ) + registered_script.assert_awaited_once_with(keys=tuple(expected_keys), args=[60], client=None) # LIT-3298: rate limits tripped at ~40M instead of 80M. async_register_script @@ -257,12 +239,8 @@ def test_async_register_script_binds_per_event_loop(namespace, monkeypatch): loop_a = asyncio.new_event_loop() loop_b = asyncio.new_event_loop() try: - result_a = loop_a.run_until_complete( - script(keys=["{k:v}:tokens"], args=[60]) - ) - result_b = loop_b.run_until_complete( - script(keys=["{k:v}:tokens"], args=[60]) - ) + result_a = loop_a.run_until_complete(script(keys=["{k:v}:tokens"], args=[60])) + result_b = loop_b.run_until_complete(script(keys=["{k:v}:tokens"], args=[60])) finally: loop_a.close() loop_b.close() @@ -275,9 +253,7 @@ def test_async_register_script_binds_per_event_loop(namespace, monkeypatch): @pytest.mark.asyncio -async def test_async_register_script_not_shared_across_namespaces( - monkeypatch, redis_no_ping -): +async def test_async_register_script_not_shared_across_namespaces(monkeypatch, redis_no_ping): """Two caches with different namespaces registering the SAME script must each run against their own client and key prefix. A content-only executor cache would let the second cache reuse the first's executor and namespace.""" @@ -293,9 +269,10 @@ async def test_async_register_script_not_shared_across_namespaces( client_b.register_script = MagicMock(return_value=reg_b) same_script = "return redis.call('GET', KEYS[1])" - with patch.object( - cache_a, "init_async_client", return_value=client_a - ), patch.object(cache_b, "init_async_client", return_value=client_b): + with ( + patch.object(cache_a, "init_async_client", return_value=client_a), + patch.object(cache_b, "init_async_client", return_value=client_b), + ): script_a = cache_a.async_register_script(same_script) script_b = cache_b.async_register_script(same_script) result_a = await script_a(keys=["k"], args=[]) @@ -307,9 +284,7 @@ async def test_async_register_script_not_shared_across_namespaces( @pytest.mark.asyncio -async def test_async_register_script_cluster_path_uses_evalsha( - monkeypatch, redis_no_ping -): +async def test_async_register_script_cluster_path_uses_evalsha(monkeypatch, redis_no_ping): """Redis Cluster exposes script_load/evalsha rather than register_script. The script is loaded once and invoked via evalsha with namespaced keys.""" monkeypatch.setenv("REDIS_HOST", "https://my-test-host") @@ -319,23 +294,17 @@ async def test_async_register_script_cluster_path_uses_evalsha( cluster_client.script_load = MagicMock(return_value="sha123") cluster_client.evalsha = AsyncMock(return_value="cluster-ok") - with patch.object( - redis_cache, "init_async_client", return_value=cluster_client - ): + with patch.object(redis_cache, "init_async_client", return_value=cluster_client): script = redis_cache.async_register_script("return 'cluster'") result = await script(keys=["{k:v}:tokens"], args=[5, 60]) assert result == "cluster-ok" cluster_client.script_load.assert_called_once_with("return 'cluster'") - cluster_client.evalsha.assert_awaited_once_with( - "sha123", 1, "ns:{k:v}:tokens", 5, 60 - ) + cluster_client.evalsha.assert_awaited_once_with("sha123", 1, "ns:{k:v}:tokens", 5, 60) @pytest.mark.asyncio -async def test_async_register_script_raises_for_unsupported_client( - monkeypatch, redis_no_ping -): +async def test_async_register_script_raises_for_unsupported_client(monkeypatch, redis_no_ping): """A client exposing neither register_script nor script_load fails loudly rather than silently returning a no-op callable.""" monkeypatch.setenv("REDIS_HOST", "https://my-test-host") @@ -350,46 +319,34 @@ async def test_async_register_script_raises_for_unsupported_client( @pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")]) @pytest.mark.asyncio -async def test_async_delete_cache_namespaces_key( - namespace, expected, monkeypatch, redis_no_ping -): +async def test_async_delete_cache_namespaces_key(namespace, expected, monkeypatch, redis_no_ping): monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache(namespace=namespace) mock_redis_instance = AsyncMock() - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): await redis_cache.async_delete_cache("k") mock_redis_instance.delete.assert_awaited_once_with(expected) @pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")]) @pytest.mark.asyncio -async def test_delete_cache_keys_namespaces_keys( - namespace, expected, monkeypatch, redis_no_ping -): +async def test_delete_cache_keys_namespaces_keys(namespace, expected, monkeypatch, redis_no_ping): monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache(namespace=namespace) mock_redis_instance = AsyncMock() - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): await redis_cache.delete_cache_keys(["k"]) mock_redis_instance.delete.assert_awaited_once_with(expected) @pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")]) @pytest.mark.asyncio -async def test_async_get_ttl_namespaces_key( - namespace, expected, monkeypatch, redis_no_ping -): +async def test_async_get_ttl_namespaces_key(namespace, expected, monkeypatch, redis_no_ping): monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache(namespace=namespace) mock_redis_instance = AsyncMock() mock_redis_instance.ttl = AsyncMock(return_value=42) - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): ttl = await redis_cache.async_get_ttl("k") assert ttl == 42 mock_redis_instance.ttl.assert_awaited_once_with(expected) @@ -397,41 +354,31 @@ async def test_async_get_ttl_namespaces_key( @pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")]) @pytest.mark.asyncio -async def test_async_lpop_namespaces_key( - namespace, expected, monkeypatch, redis_no_ping -): +async def test_async_lpop_namespaces_key(namespace, expected, monkeypatch, redis_no_ping): monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache(namespace=namespace) mock_redis_instance = AsyncMock() mock_redis_instance.lpop = AsyncMock(return_value=b"value") - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): await redis_cache.async_lpop(key="k") mock_redis_instance.lpop.assert_awaited_once_with(expected, None) @pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")]) @pytest.mark.asyncio -async def test_async_rpush_namespaces_key( - namespace, expected, monkeypatch, redis_no_ping -): +async def test_async_rpush_namespaces_key(namespace, expected, monkeypatch, redis_no_ping): monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache(namespace=namespace) mock_redis_instance = AsyncMock() mock_redis_instance.rpush = AsyncMock(return_value=1) - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): await redis_cache.async_rpush("k", ["v"]) mock_redis_instance.rpush.assert_awaited_once_with(expected, "v") @pytest.mark.parametrize("namespace, expected_match", [(None, "k*"), ("ns", "ns:k*")]) @pytest.mark.asyncio -async def test_async_scan_iter_namespaces_pattern( - namespace, expected_match, monkeypatch, redis_no_ping -): +async def test_async_scan_iter_namespaces_pattern(namespace, expected_match, monkeypatch, redis_no_ping): monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache(namespace=namespace) @@ -448,17 +395,13 @@ async def test_async_scan_iter_namespaces_pattern( mock_redis_instance = MagicMock() mock_redis_instance.scan_iter = scan_iter - with patch.object( - redis_cache, "init_async_client", return_value=mock_redis_instance - ): + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): await redis_cache.async_scan_iter(pattern="k") assert captured["match"] == expected_match @pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")]) -def test_increment_cache_namespaces_key( - namespace, expected, monkeypatch, redis_no_ping -): +def test_increment_cache_namespaces_key(namespace, expected, monkeypatch, redis_no_ping): monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache(namespace=namespace) mock_client = MagicMock() @@ -1550,3 +1493,123 @@ async def test_async_rpush_and_trim_runs_push_and_trim_in_one_transaction(monkey assert pushed_len == 4 assert rows == ["b", "c", "d"] assert pipe.queued == [("rpush", "ns:buf", "c", "d"), ("ltrim", "ns:buf", "-3", "-1")] + + +@pytest.fixture +def fake_redis_port() -> Iterator[int]: + import threading + + import fakeredis + + server = fakeredis.TcpFakeServer(("127.0.0.1", 0), server_type="redis") + worker = threading.Thread(target=server.serve_forever, daemon=True) + worker.start() + try: + yield server.server_address[1] + finally: + server.shutdown() + server.server_close() + worker.join(timeout=5) + + +@pytest.fixture +def socket_redis_cache(fake_redis_port: int, monkeypatch: pytest.MonkeyPatch) -> RedisCache: + for name in ( + "REDIS_URL", + "REDIS_HOST", + "REDIS_PORT", + "REDIS_PASSWORD", + "REDIS_CLUSTER_NODES", + "REDIS_SENTINEL_NODES", + ): + monkeypatch.delenv(name, raising=False) + return RedisCache(host="127.0.0.1", port=fake_redis_port) + + +@pytest.mark.asyncio +async def test_async_subscribe_receives_only_messages_published_after_it(socket_redis_cache: RedisCache) -> None: + from litellm.caching.redis_cache import RedisMessage + + assert await socket_redis_cache.async_publish("events", "before anyone listens") == 0 + + subscription = await socket_redis_cache.async_subscribe("events") + assert await subscription.get_message(timeout=0) is None, "the subscribe ack is not a message" + assert await socket_redis_cache.async_publish("events", "hello") == 1 + assert await socket_redis_cache.async_publish("other", b"ignored") == 0 + + assert await subscription.get_message(timeout=2) == RedisMessage(channel="events", payload=b"hello") + assert await subscription.get_message(timeout=0) is None + + await subscription.aclose() + for _ in range(50): + if await socket_redis_cache.async_publish("events", "nobody") == 0: + break + await asyncio.sleep(0.02) + else: + pytest.fail("a closed subscription still counts as a receiver") + + +@pytest.mark.asyncio +async def test_async_subscribe_covers_every_named_channel(socket_redis_cache: RedisCache) -> None: + subscription = await socket_redis_cache.async_subscribe("first", "second") + try: + await socket_redis_cache.async_publish("second", "two") + await socket_redis_cache.async_publish("first", "one") + + received = [await subscription.get_message(timeout=2) for _ in range(2)] + assert [(message.channel, message.payload) for message in received if message] == [ + ("second", b"two"), + ("first", b"one"), + ] + finally: + await subscription.aclose() + + +@pytest.mark.asyncio +async def test_pubsub_refuses_cluster_clients(fake_redis_port: int, monkeypatch: pytest.MonkeyPatch) -> None: + from redis.asyncio import RedisCluster + + for name in ( + "REDIS_URL", + "REDIS_HOST", + "REDIS_PORT", + "REDIS_PASSWORD", + "REDIS_CLUSTER_NODES", + "REDIS_SENTINEL_NODES", + ): + monkeypatch.delenv(name, raising=False) + + class _ClusterClientCache(RedisCache): + def init_async_client(self): + return RedisCluster.__new__(RedisCluster) + + cache = _ClusterClientCache(host="127.0.0.1", port=fake_redis_port) + + with pytest.raises(NotImplementedError): + await cache.async_publish("events", "hello") + with pytest.raises(NotImplementedError): + await cache.async_subscribe("events") + + +@pytest.mark.asyncio +async def test_pubsub_calls_feed_the_circuit_breaker(socket_redis_cache: RedisCache) -> None: + socket_redis_cache._circuit_breaker.record_failure() + socket_redis_cache._circuit_breaker.record_failure() + socket_redis_cache._circuit_breaker.record_failure() + socket_redis_cache._circuit_breaker.record_failure() + socket_redis_cache._circuit_breaker.record_failure() + assert socket_redis_cache._circuit_breaker.is_open() + + with pytest.raises(RedisCircuitBreakerOpenError): + await socket_redis_cache.async_publish("events", "hello") + with pytest.raises(RedisCircuitBreakerOpenError): + await socket_redis_cache.async_subscribe("events") + + +def test_connection_pool_status_reports_the_sync_pool(socket_redis_cache: RedisCache) -> None: + status = socket_redis_cache.connection_pool_status() + + assert status["max_connections"] == socket_redis_cache.redis_client.connection_pool.max_connections + assert status["connection_class"] == type( + socket_redis_cache.redis_client.connection_pool.connection_class + ).__name__ or isinstance(status["connection_class"], str) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py index e06f3bcd5cf..cb3f25714cc 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_redis_distributed_lock.py @@ -1,7 +1,21 @@ -"""Tests for the Redis lock: acquire NX/PX with a token, compare-and-delete release, namespacing.""" +"""Tests for the Redis lock: acquire NX/PX with a token, owner-only extend and release, namespacing. + +The happy paths run against a real ``redis-server`` because every operation is a Lua script; the +degrade paths use a fake ``RedisCache`` whose scripts fail. +""" + +import asyncio +import shutil +import subprocess +import time +from collections.abc import Callable, Iterator, Sequence +from pathlib import Path +from typing import Final import pytest +import redis +from litellm.caching.redis_cache import RedisCache from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_distributed_lock import ( RedisDistributedLock, ) @@ -10,117 +24,122 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_c ) -class _FakeRedis: - """Models just enough Redis to exercise SET NX (token store) and the compare-and-delete EVAL.""" - - def __init__(self, set_returns=True, exists_returns=1, raise_on=()): - self._set_returns = set_returns - self._exists_returns = exists_returns - self._raise_on = set(raise_on) - self.set_calls = [] - self.eval_calls = [] - self.expire_calls = [] - self.deleted = [] - self.store: dict = {} - - async def set(self, name, value, *, nx=False, px=None): - if "set" in self._raise_on: - raise RuntimeError("redis down") - self.set_calls.append((name, value, nx, px)) - if self._set_returns: - self.store[name] = value - return self._set_returns - - async def eval(self, script, numkeys, *keys_and_args): - if "eval" in self._raise_on: - raise RuntimeError("redis down") - self.eval_calls.append((numkeys, keys_and_args)) - key, token = keys_and_args[0], keys_and_args[1] - if len(keys_and_args) == 3: - ttl_ms = keys_and_args[2] - if self.store.get(key) == token: - self.expire_calls.append((key, ttl_ms)) - return 1 - return 0 - if self.store.get(key) == token: # compare-and-delete: only the owner deletes - del self.store[key] - self.deleted.append(key) - return 1 - return 0 - - async def exists(self, *names): - if "exists" in self._raise_on: - raise RuntimeError("redis down") - return self._exists_returns +@pytest.fixture +def redis_port(tmp_path: Path, unused_tcp_port_factory: Callable[[], int]) -> Iterator[int]: + server: Final = shutil.which("redis-server") + if server is None: + pytest.skip("redis-server is required for the lock's Lua scripts") + port: Final = unused_tcp_port_factory() + log_path: Final = tmp_path / "redis.log" + config: Final = tmp_path / "redis.conf" + config.write_text(f'bind 127.0.0.1\nport {port}\ndir "{tmp_path}"\nsave ""\nappendonly no\n') + with log_path.open("w") as log: + process: Final = subprocess.Popen((server, str(config)), stdout=log, stderr=subprocess.STDOUT) + try: + with redis.Redis(host="127.0.0.1", port=port, socket_timeout=1, socket_connect_timeout=1) as admin: + for _ in range(100): + try: + admin.ping() + break + except redis.ConnectionError: + time.sleep(0.1) + else: + pytest.fail(f"Redis did not start: {log_path.read_text()}") + yield port + finally: + process.terminate() + try: + process.wait(timeout=5) + except subprocess.TimeoutExpired: + process.kill() + process.wait() -@pytest.mark.asyncio -async def test_acquire_sets_token_with_nx_px_and_reports_acquired(): - redis = _FakeRedis(set_returns=True) - assert await RedisDistributedLock(redis).acquire("k", "tok-1", 10.0) is (LockAcquisition.ACQUIRED) - assert redis.set_calls == [("k", "tok-1", True, 10000)] # token value, NX, px in ms +@pytest.fixture +def redis_cache(redis_port: int, monkeypatch: pytest.MonkeyPatch) -> RedisCache: + for name in ( + "REDIS_URL", + "REDIS_HOST", + "REDIS_PORT", + "REDIS_PASSWORD", + "REDIS_CLUSTER_NODES", + "REDIS_SENTINEL_NODES", + ): + monkeypatch.delenv(name, raising=False) + return RedisCache(host="127.0.0.1", port=redis_port, namespace="tenant") -@pytest.mark.asyncio -async def test_acquire_reports_held_when_key_already_held(): - # redis SET NX returns None when the key exists -> another worker holds it. - assert await RedisDistributedLock(_FakeRedis(set_returns=None)).acquire("k", "tok", 10.0) is LockAcquisition.HELD +@pytest.fixture +def raw(redis_port: int) -> Iterator[redis.Redis]: + with redis.Redis(host="127.0.0.1", port=redis_port) as client: + yield client -@pytest.mark.asyncio -async def test_acquire_reports_error_on_redis_error_distinct_from_held(): - # A dead backend must be distinguishable from a busy holder so the coordinator refreshes anyway. - assert await RedisDistributedLock(_FakeRedis(raise_on=["set"])).acquire("k", "tok", 10.0) is LockAcquisition.ERROR +class _FailingScriptsCache: + """A RedisCache whose every registered script fails at run time, as a dead Redis does.""" + + def async_register_script(self, script: str) -> Callable[..., object]: + async def run(keys: Sequence[str], args: Sequence[str]) -> object: + raise ConnectionError("redis down") + + return run -@pytest.mark.asyncio -async def test_release_deletes_only_when_the_token_matches(): - # Regression: a holder whose lock PX-expired and was re-acquired by another worker must not be - # able to delete the new holder's lock. release with a stale token is a no-op. - redis = _FakeRedis() - lock = RedisDistributedLock(redis) - await lock.acquire("k", "owner-B", 10.0) # B currently holds the lock - await lock.release("k", "owner-A") # A's stale token - assert redis.deleted == [] and redis.store.get("k") == "owner-B" # B's lock survives - await lock.release("k", "owner-B") # the real owner releases - assert redis.deleted == ["k"] and "k" not in redis.store +async def test_acquire_wins_once_and_reports_held_to_the_next_caller(redis_cache: RedisCache, raw: redis.Redis) -> None: + lock = RedisDistributedLock(redis_cache) + + assert await lock.acquire("k", "tok-1", 10.0) is LockAcquisition.ACQUIRED + assert await lock.acquire("k", "tok-2", 10.0) is LockAcquisition.HELD + assert raw.get("tenant:k") == b"tok-1", "the key carries the cache namespace and the winner's token" + assert 0 < raw.pttl("tenant:k") <= 10_000 -@pytest.mark.asyncio -async def test_keys_are_namespaced_before_reaching_redis(): - redis = _FakeRedis() - lock = RedisDistributedLock(redis, namespace_key=lambda key: f"ns:{key}") - await lock.acquire("k", "tok", 10.0) - await lock.extend("k", "tok", 10.0) - await lock.release("k", "tok") - await lock.is_held("k") - assert redis.set_calls[0][0] == "ns:k" # acquire namespaced - assert redis.eval_calls[0][1][0] == "ns:k" # extend (EVAL KEYS[1]) namespaced - assert redis.eval_calls[1][1][0] == "ns:k" # release (EVAL KEYS[1]) namespaced - assert redis.deleted == ["ns:k"] +async def test_acquire_reports_error_on_redis_error_distinct_from_held() -> None: + lock = RedisDistributedLock(_FailingScriptsCache()) # pyright: ignore[reportArgumentType] # duck-typed fake + + assert await lock.acquire("k", "tok", 10.0) is LockAcquisition.ERROR -@pytest.mark.asyncio -async def test_extend_refreshes_ttl_only_when_the_token_matches(): - redis = _FakeRedis() - lock = RedisDistributedLock(redis) +async def test_release_deletes_only_when_the_token_matches(redis_cache: RedisCache) -> None: + lock = RedisDistributedLock(redis_cache) await lock.acquire("k", "owner-B", 10.0) - assert await lock.extend("k", "owner-A", 10.0) is False - assert await lock.extend("k", "owner-B", 10.0) is True - assert redis.expire_calls == [("k", "10000")] + + await lock.release("k", "owner-A") + assert await lock.is_held("k") is True, "a stale token must not delete another worker's lock" + + await lock.release("k", "owner-B") + assert await lock.is_held("k") is False -@pytest.mark.asyncio -async def test_extend_degrades_to_false_on_redis_error(): - assert await RedisDistributedLock(_FakeRedis(raise_on=["eval"])).extend("k", "tok", 10.0) is False +async def test_extend_refreshes_ttl_only_when_the_token_matches(redis_cache: RedisCache, raw: redis.Redis) -> None: + lock = RedisDistributedLock(redis_cache) + await lock.acquire("k", "owner-B", 1.0) + + assert await lock.extend("k", "owner-A", 30.0) is False + assert raw.pttl("tenant:k") <= 1_000 + assert await lock.extend("k", "owner-B", 30.0) is True + assert raw.pttl("tenant:k") > 1_000 -@pytest.mark.asyncio -async def test_is_held_reflects_exists(): - assert await RedisDistributedLock(_FakeRedis(exists_returns=1)).is_held("k") is True - assert await RedisDistributedLock(_FakeRedis(exists_returns=0)).is_held("k") is False +async def test_extend_and_release_degrade_on_redis_error() -> None: + lock = RedisDistributedLock(_FailingScriptsCache()) # pyright: ignore[reportArgumentType] # duck-typed fake + + assert await lock.extend("k", "tok", 10.0) is False + await lock.release("k", "tok") -@pytest.mark.asyncio -async def test_is_held_degrades_to_false_on_redis_error(): - assert await RedisDistributedLock(_FakeRedis(raise_on=["exists"])).is_held("k") is False +async def test_lock_expires_on_its_own_after_the_ttl(redis_cache: RedisCache) -> None: + lock = RedisDistributedLock(redis_cache) + await lock.acquire("k", "tok", 0.05) + assert await lock.is_held("k") is True + + await asyncio.sleep(0.2) + + assert await lock.is_held("k") is False + assert await lock.acquire("k", "tok-2", 10.0) is LockAcquisition.ACQUIRED + + +async def test_is_held_degrades_to_false_on_redis_error() -> None: + lock = RedisDistributedLock(_FailingScriptsCache()) # pyright: ignore[reportArgumentType] # duck-typed fake + + assert await lock.is_held("k") is False diff --git a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py index 96770ee01c4..a134501ec83 100644 --- a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py @@ -2,14 +2,14 @@ import asyncio import hashlib import json import time -from collections.abc import Iterable +from collections.abc import Awaitable, Callable, Iterable from unittest.mock import patch import pytest -from redis.asyncio import Redis import litellm.proxy.common_utils.auth_cache_invalidation_pubsub as pubsub_module from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_cache import RedisMessage from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( AUTH_CACHE_INVALIDATION_CHANNEL, @@ -20,16 +20,7 @@ from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache -class _RecordingRedisClient(Redis): - def __init__(self) -> None: - self.published: list[tuple[str, str]] = [] - - async def publish(self, channel: str, message: str) -> int: - self.published.append((channel, message)) - return 1 - - -class _WedgedPublishRedisClient(Redis): +class _WedgedPublisher: def __init__(self) -> None: self.attempted: list[str] = [] self.in_flight = 0 @@ -45,26 +36,17 @@ class _WedgedPublishRedisClient(Redis): return 1 -class _FailingPublishRedisClient(Redis): - def __init__(self) -> None: - pass - - async def publish(self, channel: str, message: str) -> int: - raise ConnectionError("redis down") - - class _QueuePubSub: - def __init__(self, initial_messages: Iterable[object] = ()) -> None: - self.queue: asyncio.Queue[object] = asyncio.Queue() + """A subscription fed from a queue of messages, standing in for RedisSubscription.""" + + def __init__(self, initial_messages: Iterable[RedisMessage] = ()) -> None: + self.queue: asyncio.Queue[RedisMessage] = asyncio.Queue() for message in initial_messages: self.queue.put_nowait(message) self.subscribed_channels: list[str] = [] self.closed = False - async def subscribe(self, *channels: str) -> None: - self.subscribed_channels.extend(channels) - - async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> object | None: + async def get_message(self, *, timeout: float | None) -> RedisMessage | None: try: return await asyncio.wait_for(self.queue.get(), timeout) except asyncio.TimeoutError: @@ -74,49 +56,70 @@ class _QueuePubSub: self.closed = True -class _ScriptedPubSubRedisClient(Redis): - def __init__(self, pubsubs: Iterable[_QueuePubSub]) -> None: - self._scripted_pubsubs = iter(pubsubs) - - def pubsub(self) -> _QueuePubSub: - return next(self._scripted_pubsubs) - - class _FakeRedisCache: - def __init__(self, client: object, namespace: str | None = None) -> None: - self._client = client + """The slice of RedisCache the module uses: namespace, async_publish and async_subscribe.""" + + def __init__( + self, + subscriptions: Iterable[_QueuePubSub] = (), + namespace: str | None = None, + publish: Callable[[str, str], Awaitable[int]] | None = None, + publish_error: Exception | None = None, + ) -> None: + self._subscriptions = iter(subscriptions) + self._publish = publish + self._publish_error = publish_error self.namespace = namespace + self.published: list[tuple[str, str]] = [] - def init_async_client(self) -> object: - return self._client + async def async_publish(self, channel: str, message: str) -> int: + if self._publish_error is not None: + raise self._publish_error + if self._publish is not None: + return await self._publish(channel, message) + self.published.append((channel, message)) + return 1 + + async def async_subscribe(self, *channels: str) -> _QueuePubSub: + subscription = next(self._subscriptions) + subscription.subscribed_channels.extend(channels) + return subscription -def _invalidation_message(cache_key: str) -> dict: - return {"type": "message", "data": json.dumps({"cache_key": cache_key}).encode()} +class _ClusterRedisCache(_FakeRedisCache): + async def async_publish(self, channel: str, message: str) -> int: + raise NotImplementedError("Redis Cluster clients have no pub/sub support") + + async def async_subscribe(self, *channels: str) -> _QueuePubSub: + raise NotImplementedError("Redis Cluster clients have no pub/sub support") + + +def _invalidation_message(cache_key: str) -> RedisMessage: + return RedisMessage(channel=AUTH_CACHE_INVALIDATION_CHANNEL, payload=json.dumps({"cache_key": cache_key}).encode()) @pytest.mark.asyncio async def test_publish_sends_cache_key_json_on_channel() -> None: - client = _RecordingRedisClient() + cache = _FakeRedisCache() with patch( "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache", - return_value=_FakeRedisCache(client=client), + return_value=cache, ): await publish_auth_cache_invalidation(cache_key="project_id:p-1") - assert client.published == [(AUTH_CACHE_INVALIDATION_CHANNEL, json.dumps({"cache_key": "project_id:p-1"}))] + assert cache.published == [(AUTH_CACHE_INVALIDATION_CHANNEL, json.dumps({"cache_key": "project_id:p-1"}))] @pytest.mark.asyncio async def test_publish_uses_namespaced_channel() -> None: - client = _RecordingRedisClient() + cache = _FakeRedisCache(namespace="ns1") with patch( "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache", - return_value=_FakeRedisCache(client=client, namespace="ns1"), + return_value=cache, ): await publish_auth_cache_invalidation(cache_key="project_id:p-1") - assert client.published[0][0] == f"ns1:{AUTH_CACHE_INVALIDATION_CHANNEL}" + assert cache.published[0][0] == f"ns1:{AUTH_CACHE_INVALIDATION_CHANNEL}" @pytest.mark.asyncio @@ -132,11 +135,33 @@ async def test_publish_noops_without_coordination_redis() -> None: async def test_publish_swallows_redis_errors() -> None: with patch( "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache", - return_value=_FakeRedisCache(client=_FailingPublishRedisClient()), + return_value=_FakeRedisCache(publish_error=ConnectionError("redis down")), ): await publish_auth_cache_invalidation(cache_key="project_id:p-1") +@pytest.mark.asyncio +async def test_publish_skips_clients_without_pubsub_support() -> None: + with patch( + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache", + return_value=_ClusterRedisCache(), + ): + await publish_auth_cache_invalidation(cache_key="project_id:p-1") + await asyncio.gather(*pubsub_module._pending_publishes) # pyright: ignore[reportPrivateUsage] # drain module-level tasks + + +@pytest.mark.asyncio +async def test_subscriber_disables_itself_without_pubsub_support() -> None: + subscriber = AuthCacheInvalidationSubscriber(redis_cache=_ClusterRedisCache(), user_api_key_cache=UserApiKeyCache()) + + subscriber.start() + task = subscriber._task + assert task is not None + await asyncio.wait_for(task, timeout=5) + + assert task.done() is True + + @pytest.mark.asyncio async def test_subscriber_deletes_local_cache_entry_on_message() -> None: """ @@ -150,7 +175,7 @@ async def test_subscriber_deletes_local_cache_entry_on_message() -> None: pubsub = _QueuePubSub(initial_messages=[_invalidation_message("project_id:p-1")]) subscriber = AuthCacheInvalidationSubscriber( - redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])), + redis_cache=_FakeRedisCache(subscriptions=[pubsub]), user_api_key_cache=cache, ) subscriber.start() @@ -180,7 +205,7 @@ async def test_subscriber_deletes_key_object_partition_entry_on_message() -> Non pubsub = _QueuePubSub(initial_messages=[_invalidation_message(hashed_token)]) subscriber = AuthCacheInvalidationSubscriber( - redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])), + redis_cache=_FakeRedisCache(subscriptions=[pubsub]), user_api_key_cache=cache, ) subscriber.start() @@ -210,7 +235,7 @@ async def test_subscriber_deletes_additional_in_memory_cache_entry_on_message() pubsub = _QueuePubSub(initial_messages=[_invalidation_message("spend:team_member:u-1:t-1")]) subscriber = AuthCacheInvalidationSubscriber( - redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])), + redis_cache=_FakeRedisCache(subscriptions=[pubsub]), user_api_key_cache=cache, additional_in_memory_caches=(spend_counter_in_memory_cache,), ) @@ -232,13 +257,13 @@ async def test_subscriber_ignores_malformed_messages() -> None: cache.in_memory_cache.set_cache("project_id:p-1", {"models": []}) subscriber = AuthCacheInvalidationSubscriber( - redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[_QueuePubSub()])), + redis_cache=_FakeRedisCache(subscriptions=[_QueuePubSub()]), user_api_key_cache=cache, ) - subscriber._apply_message({"type": "message", "data": b"not json"}) - subscriber._apply_message({"type": "message", "data": json.dumps({"other": "x"}).encode()}) - subscriber._apply_message("raw string") - subscriber._apply_message(None) + subscriber._apply_message(RedisMessage(channel=AUTH_CACHE_INVALIDATION_CHANNEL, payload=b"not json")) + subscriber._apply_message( + RedisMessage(channel=AUTH_CACHE_INVALIDATION_CHANNEL, payload=json.dumps({"other": "x"}).encode()) + ) assert cache.in_memory_cache.get_cache("project_id:p-1") is not None @@ -247,11 +272,11 @@ async def test_subscriber_ignores_malformed_messages() -> None: async def test_evict_and_broadcast_evicts_locally_and_returns_while_redis_publish_never_answers() -> None: cache = UserApiKeyCache() cache.set_cache("user-wedged", UserAPIKeyAuth(user_id="user-wedged"), model_type=UserAPIKeyAuth) - client = _WedgedPublishRedisClient() + client = _WedgedPublisher() with patch( "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache", - return_value=_FakeRedisCache(client=client), + return_value=_FakeRedisCache(publish=client.publish), ): started = time.monotonic() await evict_and_broadcast(cache_keys=("user-wedged",), user_api_key_cache=cache) @@ -270,11 +295,11 @@ async def test_publish_holds_at_most_sixteen_redis_connections_while_redis_is_we ) -> None: monkeypatch.setattr(pubsub_module, "_in_flight_publishes", asyncio.Semaphore(16)) monkeypatch.setattr(pubsub_module, "_pending_publishes", set()) - client = _WedgedPublishRedisClient() + client = _WedgedPublisher() with patch( "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache", - return_value=_FakeRedisCache(client=client), + return_value=_FakeRedisCache(publish=client.publish), ): for i in range(64): await publish_auth_cache_invalidation(cache_key=f"user-{i}") diff --git a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py index 0f64ef2b4ca..21acd8b441b 100644 --- a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py @@ -1,22 +1,22 @@ import asyncio import json import random -from typing import Callable, Coroutine, Iterable, List, Optional, Tuple +from collections.abc import Callable, Coroutine, Iterable from unittest.mock import AsyncMock, MagicMock, patch import pytest -from redis.asyncio import Redis import litellm +from litellm.caching.redis_cache import RedisMessage from litellm.proxy.common_utils.config_sync_pubsub import ( + _CONFIG_SYNCED_TABLE_NAMES, + _RESYNC_APPLIED_CONFIG_PARAM_NAMES, + _WRITE_ACTION_NAMES, CONFIG_SYNC_CHANNEL, CONFIG_SYNC_JITTER_MAX_SECONDS, CONFIG_SYNC_MIN_RESYNC_INTERVAL_SECONDS, ConfigSyncSubscriber, - _CONFIG_SYNCED_TABLE_NAMES, _PublishOnWriteActions, - _RESYNC_APPLIED_CONFIG_PARAM_NAMES, - _WRITE_ACTION_NAMES, publish_config_change, wrap_table_actions_for_config_sync, ) @@ -64,60 +64,36 @@ _EXPECTED_RESYNC_APPLIED_CONFIG_PARAM_NAMES = frozenset( _STARTUP_ONLY_CONFIG_PARAM_NAMES = ("environment_variables",) -class _RecordingRedisClient(Redis): - def __init__(self) -> None: - self.published: List[Tuple[str, str]] = [] - - async def publish(self, channel: str, message: str) -> int: - self.published.append((channel, message)) - return 1 - - -class _FailingPublishRedisClient(Redis): - def __init__(self) -> None: - pass - - async def publish(self, channel: str, message: str) -> int: - raise ConnectionError("redis down") - - -class _NotRedisClient: - def __init__(self) -> None: - self.published: List[Tuple[str, str]] = [] - - async def publish(self, channel: str, message: str) -> int: - self.published.append((channel, message)) - return 1 - - class _QueuePubSub: + """A subscription fed from a queue of payloads, standing in for RedisSubscription.""" + def __init__(self, initial_messages: Iterable[str] = ()) -> None: - self.queue: "asyncio.Queue[str]" = asyncio.Queue() + self.queue: asyncio.Queue[str] = asyncio.Queue() for message in initial_messages: self.queue.put_nowait(message) - self.subscribed_channels: List[str] = [] + self.subscribed_channels: list[str] = [] self.closed = False - async def subscribe(self, *channels: str) -> None: - self.subscribed_channels.extend(channels) - - async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Optional[str]: + async def get_message(self, *, timeout: float | None) -> RedisMessage | None: if timeout == 0: try: - return self.queue.get_nowait() + return self._message(self.queue.get_nowait()) except asyncio.QueueEmpty: return None try: - return await asyncio.wait_for(self.queue.get(), timeout) + return self._message(await asyncio.wait_for(self.queue.get(), timeout)) except asyncio.TimeoutError: return None + def _message(self, payload: str) -> RedisMessage: + return RedisMessage(channel=self.subscribed_channels[0], payload=payload.encode()) + async def aclose(self) -> None: self.closed = True class _BrokenPubSub(_QueuePubSub): - async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Optional[str]: + async def get_message(self, *, timeout: float | None) -> RedisMessage | None: raise ConnectionError("connection lost") @@ -131,11 +107,11 @@ class _EmptyPollsThenMessagePubSub(_QueuePubSub): super().__init__(initial_messages=initial_messages) self.remaining_empty_polls = empty_polls - async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Optional[str]: + async def get_message(self, *, timeout: float | None) -> RedisMessage | None: if timeout != 0 and self.remaining_empty_polls > 0: self.remaining_empty_polls -= 1 return None - return await super().get_message(ignore_subscribe_messages=ignore_subscribe_messages, timeout=timeout) + return await super().get_message(timeout=timeout) class _FakeClock: @@ -146,32 +122,44 @@ class _FakeClock: return self.now -class _ScriptedPubSubRedisClient(Redis): - def __init__(self, pubsubs: Iterable[_QueuePubSub]) -> None: - self._scripted_pubsubs = iter(pubsubs) - - def pubsub(self) -> _QueuePubSub: - return next(self._scripted_pubsubs) - - class _FakeRedisCache: - def __init__(self, client: object, namespace: Optional[str] = None) -> None: - self._client = client + """The slice of RedisCache the module uses: namespace, async_publish and async_subscribe.""" + + def __init__( + self, + subscriptions: Iterable[_QueuePubSub] = (), + namespace: str | None = None, + publish_error: Exception | None = None, + ) -> None: + self._subscriptions = iter(subscriptions) + self._publish_error = publish_error self.namespace = namespace + self.published: list[tuple[str, str]] = [] - def init_async_client(self) -> object: - return self._client + async def async_publish(self, channel: str, message: str) -> int: + if self._publish_error is not None: + raise self._publish_error + self.published.append((channel, message)) + return 1 + + async def async_subscribe(self, *channels: str) -> _QueuePubSub: + subscription = next(self._subscriptions) + subscription.subscribed_channels.extend(channels) + return subscription -class _ExplodingRedisCache: - namespace: Optional[str] = None +class _ClusterRedisCache(_FakeRedisCache): + """RedisCache over a cluster client raises NotImplementedError for both pub/sub methods.""" - def init_async_client(self) -> object: - raise ConnectionError("cannot connect") + async def async_publish(self, channel: str, message: str) -> int: + raise NotImplementedError("Redis Cluster clients have no pub/sub support") + + async def async_subscribe(self, *channels: str) -> _QueuePubSub: + raise NotImplementedError("Redis Cluster clients have no pub/sub support") def _recording_callback( - events: List[str], name: str, fired: asyncio.Event + events: list[str], name: str, fired: asyncio.Event ) -> Callable[[], Coroutine[None, None, None]]: async def callback() -> None: events.append(name) @@ -185,49 +173,44 @@ async def test_publish_noops_when_redis_cache_is_none() -> None: async def test_publish_sends_object_type_json_on_channel() -> None: - client = _RecordingRedisClient() - cache = _FakeRedisCache(client) + cache = _FakeRedisCache() await publish_config_change(redis_cache=cache, object_type="litellm_proxymodeltable") - assert len(client.published) == 1 - channel, message = client.published[0] + assert len(cache.published) == 1 + channel, message = cache.published[0] assert channel == "litellm_proxy.config_change" assert json.loads(message) == {"object_type": "litellm_proxymodeltable"} async def test_publish_uses_namespaced_channel() -> None: - client = _RecordingRedisClient() - cache = _FakeRedisCache(client, namespace="prod-eu") + cache = _FakeRedisCache(namespace="prod-eu") await publish_config_change(redis_cache=cache, object_type="litellm_credentialstable") - assert client.published[0][0] == "prod-eu:litellm_proxy.config_change" + assert cache.published[0][0] == "prod-eu:litellm_proxy.config_change" async def test_publish_swallows_redis_publish_errors() -> None: - cache = _FakeRedisCache(_FailingPublishRedisClient()) + cache = _FakeRedisCache(publish_error=ConnectionError("redis down")) await publish_config_change(redis_cache=cache, object_type="litellm_proxymodeltable") - -async def test_publish_swallows_client_init_errors() -> None: - await publish_config_change(redis_cache=_ExplodingRedisCache(), object_type="litellm_proxymodeltable") + assert cache.published == [] async def test_publish_skips_clients_without_pubsub_support() -> None: - client = _NotRedisClient() - cache = _FakeRedisCache(client) + cache = _ClusterRedisCache() await publish_config_change(redis_cache=cache, object_type="litellm_proxymodeltable") - assert client.published == [] + assert cache.published == [] async def test_subscriber_runs_injected_callbacks_in_order_on_message() -> None: pubsub = _QueuePubSub() - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) - events: List[str] = [] + cache = _FakeRedisCache(subscriptions=[pubsub]) + events: list[str] = [] fired = asyncio.Event() subscriber = ConfigSyncSubscriber( redis_cache=cache, @@ -252,8 +235,8 @@ async def test_subscriber_runs_injected_callbacks_in_order_on_message() -> None: async def test_burst_within_debounce_window_coalesces_into_one_resync() -> None: burst = [json.dumps({"object_type": "litellm_proxymodeltable"}) for _ in range(5)] pubsub = _QueuePubSub(initial_messages=burst) - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) - resyncs: List[str] = [] + cache = _FakeRedisCache(subscriptions=[pubsub]) + resyncs: list[str] = [] fired = asyncio.Event() subscriber = ConfigSyncSubscriber( redis_cache=cache, @@ -273,8 +256,8 @@ async def test_burst_within_debounce_window_coalesces_into_one_resync() -> None: async def test_subscriber_subscribes_on_namespaced_channel_and_resyncs() -> None: pubsub = _QueuePubSub() - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub]), namespace="prod-eu") - resyncs: List[str] = [] + cache = _FakeRedisCache(subscriptions=[pubsub], namespace="prod-eu") + resyncs: list[str] = [] fired = asyncio.Event() subscriber = ConfigSyncSubscriber( redis_cache=cache, @@ -299,8 +282,8 @@ class _MaxJitterRandom(random.Random): async def test_debounce_sleep_adds_jitter_from_injected_rng() -> None: pubsub = _QueuePubSub(initial_messages=[json.dumps({"object_type": "litellm_proxymodeltable"})]) - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) - sleeps: List[float] = [] + cache = _FakeRedisCache(subscriptions=[pubsub]) + sleeps: list[float] = [] fired = asyncio.Event() async def recording_sleep(seconds: float) -> None: @@ -332,7 +315,7 @@ def test_default_min_resync_interval_caps_reload_rate() -> None: def _throttled_subscriber( cache: object, - events: List[str], + events: list[str], fired: asyncio.Event, clock: _FakeClock, min_resync_interval_seconds: float = 10.0, @@ -358,8 +341,8 @@ def _throttled_subscriber( async def test_resync_arriving_inside_min_interval_waits_out_the_remainder() -> None: pubsub = _QueuePubSub() - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) - events: List[str] = [] + cache = _FakeRedisCache(subscriptions=[pubsub]) + events: list[str] = [] fired = asyncio.Event() clock = _FakeClock() subscriber = _throttled_subscriber(cache=cache, events=events, fired=fired, clock=clock) @@ -378,8 +361,8 @@ async def test_resync_arriving_inside_min_interval_waits_out_the_remainder() -> async def test_resync_after_min_interval_elapsed_is_not_throttled() -> None: pubsub = _QueuePubSub() - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) - events: List[str] = [] + cache = _FakeRedisCache(subscriptions=[pubsub]) + events: list[str] = [] fired = asyncio.Event() clock = _FakeClock() subscriber = _throttled_subscriber(cache=cache, events=events, fired=fired, clock=clock) @@ -398,8 +381,8 @@ async def test_resync_after_min_interval_elapsed_is_not_throttled() -> None: async def test_writes_during_the_throttle_wait_collapse_into_the_next_resync() -> None: pubsub = _QueuePubSub() - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) - events: List[str] = [] + cache = _FakeRedisCache(subscriptions=[pubsub]) + events: list[str] = [] fired = asyncio.Event() clock = _FakeClock() subscriber = _throttled_subscriber(cache=cache, events=events, fired=fired, clock=clock) @@ -420,8 +403,8 @@ async def test_writes_during_the_throttle_wait_collapse_into_the_next_resync() - async def test_polls_without_messages_do_not_trigger_resyncs() -> None: pubsub = _EmptyPollsThenMessagePubSub(empty_polls=3, initial_messages=["change"]) - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) - resyncs: List[str] = [] + cache = _FakeRedisCache(subscriptions=[pubsub]) + resyncs: list[str] = [] fired = asyncio.Event() subscriber = ConfigSyncSubscriber( redis_cache=cache, @@ -442,8 +425,8 @@ async def test_polls_without_messages_do_not_trigger_resyncs() -> None: async def test_failing_pubsub_close_still_reconnects() -> None: broken = _CloseFailingBrokenPubSub() healthy = _QueuePubSub(initial_messages=["change"]) - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([broken, healthy])) - resyncs: List[str] = [] + cache = _FakeRedisCache(subscriptions=[broken, healthy]) + resyncs: list[str] = [] fired = asyncio.Event() subscriber = ConfigSyncSubscriber( redis_cache=cache, @@ -464,7 +447,7 @@ async def test_failing_pubsub_close_still_reconnects() -> None: async def test_second_start_does_not_open_a_second_subscription() -> None: pubsub = _QueuePubSub() - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) + cache = _FakeRedisCache(subscriptions=[pubsub]) subscriber = ConfigSyncSubscriber(redis_cache=cache, resync_callbacks=(), debounce_seconds=0.01) subscriber.start() @@ -481,8 +464,8 @@ async def test_second_start_does_not_open_a_second_subscription() -> None: async def test_redis_error_leads_to_backoff_and_resubscribe() -> None: broken = _BrokenPubSub() healthy = _QueuePubSub(initial_messages=[json.dumps({"object_type": "litellm_credentialstable"})]) - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([broken, healthy])) - resyncs: List[str] = [] + cache = _FakeRedisCache(subscriptions=[broken, healthy]) + resyncs: list[str] = [] fired = asyncio.Event() subscriber = ConfigSyncSubscriber( redis_cache=cache, @@ -508,8 +491,8 @@ async def test_redis_error_leads_to_backoff_and_resubscribe() -> None: async def test_failing_resync_callback_does_not_kill_subscriber() -> None: pubsub = _QueuePubSub() - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) - resyncs: List[str] = [] + cache = _FakeRedisCache(subscriptions=[pubsub]) + resyncs: list[str] = [] fired = asyncio.Event() async def failing_callback() -> None: @@ -536,7 +519,7 @@ async def test_failing_resync_callback_does_not_kill_subscriber() -> None: async def test_stop_cancels_subscriber_cleanly() -> None: pubsub = _QueuePubSub() - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([pubsub])) + cache = _FakeRedisCache(subscriptions=[pubsub]) subscriber = ConfigSyncSubscriber(redis_cache=cache, resync_callbacks=(), debounce_seconds=0.01) subscriber.start() @@ -552,15 +535,15 @@ async def test_stop_cancels_subscriber_cleanly() -> None: async def test_stop_before_start_is_a_noop() -> None: - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([])) + cache = _FakeRedisCache(subscriptions=[]) subscriber = ConfigSyncSubscriber(redis_cache=cache, resync_callbacks=()) await subscriber.stop() async def test_subscriber_exits_without_callbacks_when_client_lacks_pubsub() -> None: - cache = _FakeRedisCache(_NotRedisClient()) - resyncs: List[str] = [] + cache = _ClusterRedisCache() + resyncs: list[str] = [] subscriber = ConfigSyncSubscriber( redis_cache=cache, resync_callbacks=(_recording_callback(resyncs, "resync", asyncio.Event()),), @@ -575,7 +558,7 @@ async def test_subscriber_exits_without_callbacks_when_client_lacks_pubsub() -> class _FakeTableActions: - def __init__(self, calls: List[Tuple[str, str]]) -> None: + def __init__(self, calls: list[tuple[str, str]]) -> None: self._calls = calls async def create(self, **kwargs: object) -> object: @@ -588,7 +571,7 @@ class _FakeTableActions: class _AllWritesTableActions: - def __init__(self, calls: List[str]) -> None: + def __init__(self, calls: list[str]) -> None: self._calls = calls def __getattr__(self, name: str) -> Callable[..., Coroutine[None, None, str]]: @@ -599,7 +582,7 @@ class _AllWritesTableActions: return action -def _recording_publish(calls: List[Tuple[str, str]]) -> Callable[[str], Coroutine[None, None, None]]: +def _recording_publish(calls: list[tuple[str, str]]) -> Callable[[str], Coroutine[None, None, None]]: async def publish(object_type: str) -> None: calls.append(("publish", object_type)) @@ -615,7 +598,7 @@ def test_wrapper_passes_through_unsynced_tables() -> None: async def test_wrapper_publishes_table_name_after_write() -> None: - calls: List[Tuple[str, str]] = [] + calls: list[tuple[str, str]] = [] wrapped = wrap_table_actions_for_config_sync( actions=_FakeTableActions(calls), table_name="litellm_proxymodeltable", @@ -629,7 +612,7 @@ async def test_wrapper_publishes_table_name_after_write() -> None: async def test_wrapper_does_not_publish_on_reads() -> None: - calls: List[Tuple[str, str]] = [] + calls: list[tuple[str, str]] = [] wrapped = wrap_table_actions_for_config_sync( actions=_FakeTableActions(calls), table_name="litellm_proxymodeltable", @@ -660,8 +643,8 @@ def test_tool_telemetry_table_writes_pass_through_unwrapped() -> None: @pytest.mark.parametrize("action_name", _EXPECTED_WRITE_ACTION_NAMES) async def test_wrapper_publishes_for_every_write_action(action_name: str) -> None: - write_calls: List[str] = [] - publish_calls: List[Tuple[str, str]] = [] + write_calls: list[str] = [] + publish_calls: list[tuple[str, str]] = [] wrapped = wrap_table_actions_for_config_sync( actions=_AllWritesTableActions(write_calls), table_name="litellm_guardrailstable", @@ -680,7 +663,7 @@ async def test_model_repository_write_publishes_via_live_coordination_cache() -> from litellm.proxy.proxy_server import _set_redis_usage_cache from litellm.repositories.model_repository import ModelRepository - client = _RecordingRedisClient() + client = _FakeRedisCache() prisma_client = MagicMock() prisma_client.db.litellm_proxymodeltable.update = AsyncMock(return_value={"model_id": "m-1"}) repository = ModelRepository(prisma_client) @@ -688,7 +671,7 @@ async def test_model_repository_write_publishes_via_live_coordination_cache() -> assert isinstance(table, _PublishOnWriteActions) previous_cache = proxy_server.redis_usage_cache - _set_redis_usage_cache(_FakeRedisCache(client)) + _set_redis_usage_cache(client) try: await table.update(where={"model_id": "m-1"}, data={"model_name": "gpt-5.2"}) finally: @@ -708,14 +691,14 @@ async def test_ui_settings_write_publishes_via_live_coordination_cache() -> None from litellm.proxy.proxy_server import _set_redis_usage_cache from litellm.repositories.table_repositories import UISettingsRepository - client = _RecordingRedisClient() + client = _FakeRedisCache() prisma_client = MagicMock() prisma_client.db.litellm_uisettings.upsert = AsyncMock(return_value={"id": "ui_settings"}) table = UISettingsRepository(prisma_client).table assert isinstance(table, _PublishOnWriteActions) previous_cache = proxy_server.redis_usage_cache - _set_redis_usage_cache(_FakeRedisCache(client)) + _set_redis_usage_cache(client) try: await table.upsert( where={"id": "ui_settings"}, @@ -730,14 +713,14 @@ async def test_ui_settings_write_publishes_via_live_coordination_cache() -> None assert json.loads(message) == {"object_type": "litellm_uisettings"} -async def _publish_calls_for_invalidated_param(param_name: str) -> List[Tuple[str, str]]: +async def _publish_calls_for_invalidated_param(param_name: str) -> list[tuple[str, str]]: from litellm.proxy import proxy_server from litellm.proxy.proxy_server import _set_redis_usage_cache from litellm.proxy.utils import invalidate_config_param - client = _RecordingRedisClient() + client = _FakeRedisCache() previous_cache = proxy_server.redis_usage_cache - _set_redis_usage_cache(_FakeRedisCache(client)) + _set_redis_usage_cache(client) try: await invalidate_config_param(param_name) finally: @@ -771,9 +754,9 @@ async def test_evict_config_param_does_not_publish() -> None: from litellm.proxy.proxy_server import _set_redis_usage_cache from litellm.proxy.utils import evict_config_param - client = _RecordingRedisClient() + client = _FakeRedisCache() previous_cache = proxy_server.redis_usage_cache - _set_redis_usage_cache(_FakeRedisCache(client)) + _set_redis_usage_cache(client) try: await evict_config_param("model_cost_map_reload_config") finally: @@ -803,10 +786,10 @@ async def test_model_cost_map_reload_does_not_publish_config_change() -> None: litellm_config_cache.flush_cache() prisma_client = _reload_config_prisma_client() - client = _RecordingRedisClient() + client = _FakeRedisCache() previous_cache = proxy_server.redis_usage_cache original_model_cost = litellm.model_cost.copy() - _set_redis_usage_cache(_FakeRedisCache(client)) + _set_redis_usage_cache(client) try: from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded @@ -833,9 +816,9 @@ async def test_anthropic_beta_headers_reload_does_not_publish_config_change() -> litellm_config_cache.flush_cache() prisma_client = _reload_config_prisma_client() - client = _RecordingRedisClient() + client = _FakeRedisCache() previous_cache = proxy_server.redis_usage_cache - _set_redis_usage_cache(_FakeRedisCache(client)) + _set_redis_usage_cache(client) try: with patch("litellm.anthropic_beta_headers_manager.reload_beta_headers_config") as mock_reload: mock_reload.return_value = {} @@ -855,11 +838,11 @@ class _StopFailingSubscriber(ConfigSyncSubscriber): async def test_proxy_config_subscriber_resyncs_deployments_only() -> None: from litellm.proxy.proxy_server import ProxyConfig - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([_QueuePubSub()])) + cache = _FakeRedisCache(subscriptions=[_QueuePubSub()]) config = ProxyConfig() prisma_client = MagicMock() proxy_logging_obj = MagicMock() - calls: List[Tuple[str, object, object]] = [] + calls: list[tuple[str, object, object]] = [] async def fake_add_deployment(prisma_client: object, proxy_logging_obj: object) -> None: calls.append(("add_deployment", prisma_client, proxy_logging_obj)) @@ -902,7 +885,7 @@ async def test_proxy_config_does_not_start_subscriber_without_coordination_redis async def test_proxy_config_keeps_the_first_subscriber_on_repeat_start() -> None: from litellm.proxy.proxy_server import ProxyConfig - cache = _FakeRedisCache(_ScriptedPubSubRedisClient([_QueuePubSub()])) + cache = _FakeRedisCache(subscriptions=[_QueuePubSub()]) config = ProxyConfig() config.start_config_sync_subscriber(prisma_client=MagicMock(), proxy_logging_obj=MagicMock(), redis_cache=cache) @@ -920,7 +903,7 @@ async def test_proxy_config_shutdown_survives_a_failing_subscriber_stop() -> Non config = ProxyConfig() config.config_sync_subscriber = _StopFailingSubscriber( - redis_cache=_FakeRedisCache(_ScriptedPubSubRedisClient([])), + redis_cache=_FakeRedisCache(subscriptions=[]), resync_callbacks=(), ) diff --git a/tests/test_litellm/proxy/common_utils/test_debug_utils.py b/tests/test_litellm/proxy/common_utils/test_debug_utils.py index a64c94c2991..c40c323025e 100644 --- a/tests/test_litellm/proxy/common_utils/test_debug_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_debug_utils.py @@ -146,3 +146,28 @@ def test_debug_report_returns_what_the_bug_report_link_carries_and_nothing_from_ "model_list[*].provider = [azure]", ] assert not any(hostile in response.text for hostile in HOSTILE_STRINGS), response.text + + +def test_cache_stats_report_the_redis_pool_through_the_cache_object() -> None: + from types import SimpleNamespace + + from litellm.caching.dual_cache import DualCache + from litellm.proxy.common_utils.debug_utils import _get_cache_memory_stats + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + class _PooledRedisCache: + def connection_pool_status(self) -> Mapping[str, object]: + return {"max_connections": 7, "connection_class": "Connection"} + + stats = _get_cache_memory_stats( + user_api_key_cache=UserApiKeyCache(), + llm_router=None, + proxy_logging_obj=SimpleNamespace(internal_usage_cache=SimpleNamespace(dual_cache=DualCache())), + redis_usage_cache=_PooledRedisCache(), + ) + + assert stats["redis_usage_cache"] == { + "enabled": True, + "cache_type": "_PooledRedisCache", + "connection_pool": {"max_connections": 7, "connection_class": "Connection"}, + }