feat(caching): give RedisCache a pub/sub and lease surface and drop raw client reaches

Add async_publish and async_subscribe to the Python RedisCache, returning a
RedisSubscription that yields RedisMessage values and skips redis-py's
subscribe acks, plus connection_pool_status for the debug endpoint. Cluster
clients raise NotImplementedError, matching the existing skip behaviour.

Move the config sync and auth cache invalidation modules onto that surface
and rewrite the MCP refresh lock onto async_register_script, so no proxy code
reaches into init_async_client or redis_client any more. This defines the
surface a native RedisCache has to match.

On the Rust side add the PubSubCache and MessageStream capability traits and
implement them in cache-redis with a dedicated subscription connection driven
by a reader thread, tested against an in-process RESP server over a real
socket.
This commit is contained in:
Yujong Lee 2026-09-24 09:19:14 -07:00
parent cfa2830bde
commit 9b747bf337
20 changed files with 1251 additions and 547 deletions

View file

@ -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

View file

@ -21,7 +21,10 @@ const REDIS_POOL_SIZE: u32 = 16;
#[allow(private_interfaces)]
pub enum Connections<C> {
Pool(r2d2::Pool<ConnectionManager>),
Pool {
pool: r2d2::Pool<ConnectionManager>,
client: Box<redis::Client>,
},
Cluster(r2d2::Pool<ClusterConnectionManager>),
Fixed(Mutex<C>),
}
@ -35,7 +38,7 @@ where
operation: impl FnOnce(&mut ConnectionRef<'_>) -> Result<T, Error>,
) -> Result<T, Error> {
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<Self, Error> {
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<redis::Connection, Error> {
match self {
Self::Pool { client, .. } => client.get_connection().map_err(|_| Error::Unavailable),
Self::Cluster(_) | Self::Fixed(_) => Err(Error::UnsupportedOperation),
}
}
}
fn pool<M: r2d2::ManageConnection>(manager: M) -> Result<r2d2::Pool<M>, Error> {
@ -119,14 +137,6 @@ pub(crate) struct PooledConnection<C> {
/// 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<Self, Error> {
redis::Client::open(url)
.map(Self)
.map_err(|_| Error::Unavailable)
}
}
impl r2d2::ManageConnection for ConnectionManager {
type Connection = PooledConnection<redis::Connection>;
type Error = redis::RedisError;

View file

@ -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};

View file

@ -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<Result<Message, Error>>,
stop: Arc<AtomicBool>,
reader: Option<thread::JoinHandle<()>>,
}
impl MessageStream for RedisSubscription {
async fn next_message(&mut self, timeout: Option<Duration>) -> Result<Option<Message>, 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<S, C> PubSubCache for RedisCache<S, C>
where
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
type Subscription = RedisSubscription;
async fn async_publish(&self, channel: &str, payload: &[u8]) -> Result<usize, Error> {
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::<usize>(connection)
.map_err(|_| Error::Unavailable)
})
.await
}
async fn async_subscribe(&self, channels: &[String]) -> Result<Self::Subscription, Error> {
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<Result<(), Error>>,
sender: &mpsc::Sender<Result<Message, Error>>,
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;
}
}
}
}

View file

@ -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<serde_json::Value>;
fn live(server: &PubSubServer) -> RedisCache<Json> {
RedisCache::new(&server.url(), None, JsonCodec::new()).unwrap()
}
fn channels(names: &[&str]) -> Vec<String> {
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<Json> = RedisCache::new(&url, None, JsonCodec::new()).unwrap();
assert_eq!(
cache
.async_subscribe(&channels(&["events"]))
.await
.map(|_| ())
.unwrap_err(),
Error::Unavailable
);
}

View file

@ -1,5 +1,7 @@
#![allow(dead_code)]
pub mod server;
use std::{
collections::BTreeMap,
time::{Duration, SystemTime, UNIX_EPOCH},

View file

@ -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<String, Vec<(u64, TcpStream)>>,
unsubscribed: Vec<String>,
}
/// 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<Mutex<Registry>>,
shutdown: Arc<AtomicBool>,
acceptor: Option<thread::JoinHandle<()>>,
}
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(&registry);
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(&registry);
thread::spawn(move || serve(stream, id, &registry));
}
}
});
Self {
address,
registry,
shutdown,
acceptor: Some(acceptor),
}
}
pub fn url(&self) -> String {
format!("redis://{}", self.address)
}
pub fn unsubscribed(&self) -> Vec<String> {
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<Registry>) {
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<u8>],
stream: &TcpStream,
id: u64,
registry: &Mutex<Registry>,
) -> Vec<u8> {
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(&registry, id);
array(&[bulk(b"subscribe"), bulk(channel.as_bytes()), integer(count)])
})
.collect()
}
"UNSUBSCRIBE" | "PUNSUBSCRIBE" => {
let mut registry = registry.lock().unwrap();
let named: Vec<String> = command[1..].iter().map(|channel| text(channel)).collect();
let channels: Vec<String> = 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(&registry, 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<Vec<u8>>, 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<usize> {
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<u8> {
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<u8> {
format!(":{value}\r\n").into_bytes()
}
fn array(items: &[Vec<u8>]) -> Vec<u8> {
let mut out = format!("*{}\r\n", items.len()).into_bytes();
for item in items {
out.extend_from_slice(item);
}
out
}

View file

@ -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<u8>,
}
/// 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<Duration>,
) -> impl Future<Output = Result<Option<Message>, Error>> + Send;
fn close(self) -> impl Future<Output = Result<(), Error>> + 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<Output = Result<usize, Error>> + Send;
fn async_subscribe(
&self,
channels: &[String],
) -> impl Future<Output = Result<Self::Subscription, Error>> + Send;
}

View file

@ -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::{

View file

@ -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,

View file

@ -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

View file

@ -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))

View file

@ -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)

View file

@ -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)

View file

@ -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}

View file

@ -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)

View file

@ -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

View file

@ -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}")

View file

@ -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=(),
)

View file

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