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"}, + }