diff --git a/litellm-rust/crates/cache-redis/src/cache.rs b/litellm-rust/crates/cache-redis/src/cache.rs index d2e94d75318..e2e2656fcbb 100644 --- a/litellm-rust/crates/cache-redis/src/cache.rs +++ b/litellm-rust/crates/cache-redis/src/cache.rs @@ -88,8 +88,6 @@ impl RedisCache { Self::connect(url, &RedisTopology::Standalone, default_ttl, codec) } - /// The URL carries credentials, database, protocol and TLS mode. For a cluster topology its - /// address is replaced by each startup node; slot discovery then finds the remaining nodes. pub fn connect( url: &str, topology: &RedisTopology, @@ -333,7 +331,7 @@ where async fn test_connection(&self) -> Result { match Self::run_blocking(Arc::clone(&self.connections), |connection| { - Ok(match redis::cmd("PING").query::(connection) { + Ok(match connection.ping() { Ok(_) => CacheConnectionResult { status: CacheConnectionStatus::Success, message: "Redis cache connection test successful".into(), diff --git a/litellm-rust/crates/cache-redis/src/cache/connection.rs b/litellm-rust/crates/cache-redis/src/cache/connection.rs index dc7afa296f2..1834f1d94e5 100644 --- a/litellm-rust/crates/cache-redis/src/cache/connection.rs +++ b/litellm-rust/crates/cache-redis/src/cache/connection.rs @@ -113,7 +113,7 @@ impl r2d2::ManageConnection for ClusterConnectionManager { } fn has_broken(&self, connection: &mut Self::Connection) -> bool { - connection.failed || !connection.connection.check_connection() + connection.failed || !redis::ConnectionLike::is_open(&connection.connection) } } @@ -252,6 +252,24 @@ impl ConnectionRef<'_> { Ok(()) } + pub(crate) fn ping(&mut self) -> Result { + let command = redis::cmd("PING"); + match self { + Self::Node(connection) => command + .query::(*connection) + .map(|response| response == "PONG"), + Self::Cluster(connection) => connection + .route_command( + &command, + RoutingInfo::MultiNode(( + MultipleNodeRoutingInfo::AllNodes, + Some(ResponsePolicy::AllSucceeded), + )), + ) + .map(|_| true), + } + } + pub(crate) fn node_text(&mut self, command: &redis::Cmd) -> Result { match self { Self::Node(connection) => command.query(*connection).map_err(|_| Error::Unavailable), diff --git a/litellm-rust/crates/cache-redis/src/cache/operations.rs b/litellm-rust/crates/cache-redis/src/cache/operations.rs index c6fc5eb27e6..9a7023338bf 100644 --- a/litellm-rust/crates/cache-redis/src/cache/operations.rs +++ b/litellm-rust/crates/cache-redis/src/cache/operations.rs @@ -183,20 +183,13 @@ where } pub fn sync_ping(&self) -> Result { - self.connections.execute(|connection| { - redis::cmd("PING") - .query::(connection) - .map(|response| response == "PONG") - .map_err(|_| Error::Unavailable) - }) + self.connections + .execute(|connection| connection.ping().map_err(|_| Error::Unavailable)) } pub async fn ping(&self) -> Result { Self::run_blocking(Arc::clone(&self.connections), |connection| { - redis::cmd("PING") - .query::(connection) - .map(|response| response == "PONG") - .map_err(|_| Error::Unavailable) + connection.ping().map_err(|_| Error::Unavailable) }) .await } diff --git a/litellm-rust/crates/cache-redis/tests/cluster.rs b/litellm-rust/crates/cache-redis/tests/cluster.rs index a0f51416362..2c3fc818b66 100644 --- a/litellm-rust/crates/cache-redis/tests/cluster.rs +++ b/litellm-rust/crates/cache-redis/tests/cluster.rs @@ -330,6 +330,53 @@ async fn scan_and_scoped_flush_cover_every_primary() { other.async_flush_cache().await.unwrap(); } +fn ping_calls_per_node(startup: &redis::Client) -> Vec<(String, u64)> { + let mut connection = startup.get_connection().unwrap(); + let nodes: String = redis::cmd("CLUSTER") + .arg("NODES") + .query(&mut connection) + .unwrap(); + let mut counts: Vec<(String, u64)> = nodes + .lines() + .map(|line| { + let address = line.split_whitespace().nth(1).unwrap(); + let address = address.split('@').next().unwrap(); + let mut node = redis::Client::open(format!("redis://{address}")) + .unwrap() + .get_connection() + .unwrap(); + let stats: String = redis::cmd("INFO") + .arg("commandstats") + .query(&mut node) + .unwrap(); + let calls = stats + .lines() + .find_map(|stat| stat.strip_prefix("cmdstat_ping:calls=")) + .and_then(|rest| rest.split(',').next()) + .map_or(0, |calls| calls.parse().unwrap()); + (address.to_string(), calls) + }) + .collect(); + counts.sort(); + counts +} + +#[tokio::test] +async fn ping_reaches_every_node() { + let cache = cluster_or_skip!("ping"); + let startup = redis::Client::open(cluster_url()).unwrap(); + let before = ping_calls_per_node(&startup); + assert!(before.len() >= 2, "{before:?}"); + assert!(cache.ping().await.unwrap()); + let after = ping_calls_per_node(&startup); + for ((node, calls_before), (_, calls_after)) in before.iter().zip(&after) { + assert!(calls_after > calls_before, "{node} was not pinged"); + } + assert!(cache.sync_ping().unwrap()); + let result = cache.test_connection().await.unwrap(); + assert_eq!(result.status, CacheConnectionStatus::Success); +} + #[tokio::test] async fn counters_claims_scripts_and_sets_work_on_the_cluster() { let Some(counter) = counter_cache("counter") else {