diff --git a/crates/dbx-core/src/db/redis_driver.rs b/crates/dbx-core/src/db/redis_driver.rs index 1a76625b3..fa771021b 100644 --- a/crates/dbx-core/src/db/redis_driver.rs +++ b/crates/dbx-core/src/db/redis_driver.rs @@ -9,6 +9,7 @@ use redis::{ TlsMode, Value as RedisRawValue, }; use serde::{Deserialize, Serialize}; +use std::{collections::HashMap, future::Future, time::Duration, time::Instant}; use tokio::sync::{Mutex, MutexGuard}; use super::json_value_for_js; @@ -20,6 +21,10 @@ const DEFAULT_REDIS_DATABASES: u32 = 16; const CLUSTER_CURSOR_NODE_BITS: u64 = 16; const CLUSTER_CURSOR_NODE_MASK: u64 = (1 << CLUSTER_CURSOR_NODE_BITS) - 1; const CLUSTER_CURSOR_SCAN_MASK: u64 = (1 << (64 - CLUSTER_CURSOR_NODE_BITS)) - 1; +const CLUSTER_SCAN_SESSION_LIMIT: usize = 128; +const CLUSTER_SCAN_SESSION_TTL: Duration = Duration::from_secs(5 * 60); +const INVALID_CLUSTER_SCAN_CURSOR_ERROR: &str = "Redis cluster scan cursor is invalid or expired"; +const MAX_SAFE_INTEGER_CURSOR: u64 = (1 << 53) - 1; #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RedisDatabaseInfo { @@ -220,6 +225,70 @@ pub struct RedisClusterPool { pub tls_insecure: bool, pub username: String, pub password: String, + scan_sessions: Box>, +} + +#[derive(Debug, Clone)] +struct RedisClusterKeyScanSession { + master_nodes: Vec, + node_index: usize, + node_cursor: u64, + pattern: String, + last_used: Instant, +} + +impl RedisClusterKeyScanSession { + fn new(master_nodes: Vec, pattern: &str) -> Self { + Self { master_nodes, node_index: 0, node_cursor: 0, pattern: pattern.to_string(), last_used: Instant::now() } + } +} + +#[derive(Debug)] +struct RedisClusterScanSessions { + next_cursor: u64, + entries: HashMap, +} + +impl Default for RedisClusterScanSessions { + fn default() -> Self { + Self { next_cursor: 1, entries: HashMap::new() } + } +} + +impl RedisClusterScanSessions { + fn take(&mut self, cursor: u64) -> Option { + self.remove_expired(); + self.entries.remove(&cursor) + } + + fn insert(&mut self, cursor: Option, mut session: RedisClusterKeyScanSession) -> u64 { + self.remove_expired(); + while self.entries.len() >= CLUSTER_SCAN_SESSION_LIMIT { + let Some(oldest) = self.entries.iter().min_by_key(|(_, entry)| entry.last_used).map(|(id, _)| *id) else { + break; + }; + self.entries.remove(&oldest); + } + + let cursor = cursor.filter(|value| *value > 0).unwrap_or_else(|| self.next_available_cursor()); + session.last_used = Instant::now(); + self.entries.insert(cursor, session); + cursor + } + + fn remove_expired(&mut self) { + self.entries.retain(|_, entry| entry.last_used.elapsed() < CLUSTER_SCAN_SESSION_TTL); + } + + fn next_available_cursor(&mut self) -> u64 { + loop { + let cursor = self.next_cursor; + self.next_cursor = if cursor >= MAX_SAFE_INTEGER_CURSOR { 1 } else { cursor + 1 }; + if cursor != 0 && !self.entries.contains_key(&cursor) { + return cursor; + } + } + } } #[derive(Debug, Clone, PartialEq, Eq)] @@ -470,6 +539,7 @@ pub async fn connect_cluster(config: &ConnectionConfig) -> Result { @@ -537,6 +607,7 @@ pub async fn connect_routed_cluster( tls_insecure: config.redis_tls_insecure(), username: auth.username, password: auth.password, + scan_sessions: Box::new(Mutex::new(RedisClusterScanSessions::default())), }; if pool.node_routes.is_empty() { pool.node_routes = identity_routes(&unique_master_nodes(&pool.slot_ranges)); @@ -942,42 +1013,170 @@ pub async fn scan_cluster_keys_page_with_options( count: usize, include_types: bool, ) -> Result { - let master_nodes = cluster_master_nodes(pool).await?; - if master_nodes.is_empty() { - return Ok(RedisScanResult { cursor: 0, keys: Vec::new(), total_keys: 0 }); + scan_cluster_keys_batch(pool, cursor, pattern, count, 1, include_types).await +} + +pub async fn scan_cluster_keys_batch( + pool: &RedisClusterPool, + cursor: u64, + pattern: &str, + count: usize, + max_iterations: usize, + include_types: bool, +) -> Result { + scan_cluster_keys_batch_with( + &pool.scan_sessions, + || async { cluster_master_nodes(pool).await }, + |endpoint| async move { connect_cluster_node(pool, &endpoint).await }, + cursor, + pattern, + count, + max_iterations, + include_types, + ) + .await +} + +async fn scan_cluster_keys_batch_with( + sessions: &Mutex, + mut discover_master_nodes: Discover, + connect_node: Connect, + cursor: u64, + pattern: &str, + count: usize, + max_iterations: usize, + include_types: bool, +) -> Result +where + C: ConnectionLike + Send + Sync + Unpin, + Discover: FnMut() -> DiscoverFuture, + DiscoverFuture: Future, String>>, + Connect: FnMut(RedisNodeEndpoint) -> ConnectFuture, + ConnectFuture: Future>, +{ + let (mut session, can_continue) = if cursor == 0 { + let master_nodes = canonical_cluster_master_nodes(discover_master_nodes().await?); + if master_nodes.is_empty() { + return Ok(RedisScanResult { cursor: 0, keys: Vec::new(), total_keys: 0 }); + } + (RedisClusterKeyScanSession::new(master_nodes, pattern), false) + } else { + let previous_session = sessions.lock().await.take(cursor); + match previous_session { + Some(session) if session.pattern == pattern => (session, true), + Some(session) => { + sessions.lock().await.insert(Some(cursor), session); + return Err(INVALID_CLUSTER_SCAN_CURSOR_ERROR.to_string()); + } + None => return Err(INVALID_CLUSTER_SCAN_CURSOR_ERROR.to_string()), + } + }; + let include_total_keys = !can_continue; + let retry_session = can_continue.then(|| session.clone()); + + let batch = match scan_cluster_keys_batch_on_session( + &mut session, + connect_node, + pattern, + count, + max_iterations, + include_types, + include_total_keys, + ) + .await + { + Ok(batch) => batch, + Err(error) => { + if let Some(retry_session) = retry_session { + // A failed batch may have advanced across nodes without returning its keys, so retries must resume + // from the request's original position rather than the partially mutated session. + sessions.lock().await.insert(Some(cursor), retry_session); + } + return Err(error); + } + }; + + if batch.complete { + return Ok(RedisScanResult { cursor: 0, keys: batch.keys, total_keys: batch.total_keys }); } - let (mut node_index, node_cursor) = decode_cluster_cursor(cursor); - if node_index >= master_nodes.len() { - node_index = 0; + let session_cursor = sessions.lock().await.insert(can_continue.then_some(cursor), session); + Ok(RedisScanResult { cursor: session_cursor, keys: batch.keys, total_keys: batch.total_keys }) +} + +struct RedisClusterKeyScanBatch { + keys: Vec, + total_keys: u64, + complete: bool, +} + +async fn scan_cluster_keys_batch_on_session( + session: &mut RedisClusterKeyScanSession, + mut connect_node: Connect, + pattern: &str, + count: usize, + max_iterations: usize, + include_types: bool, + include_total_keys: bool, +) -> Result +where + C: ConnectionLike + Send + Sync + Unpin, + Connect: FnMut(RedisNodeEndpoint) -> ConnectFuture, + ConnectFuture: Future>, +{ + if session.master_nodes.is_empty() || session.node_index >= session.master_nodes.len() { + return Ok(RedisClusterKeyScanBatch { keys: Vec::new(), total_keys: 0, complete: true }); } - let total_keys = cluster_total_keys(pool, &master_nodes).await; - for index in node_index..master_nodes.len() { - let endpoint = &master_nodes[index]; - let mut con = connect_cluster_node(pool, endpoint).await?; - let current_cursor = if index == node_index { node_cursor } else { 0 }; - let result = scan_keys_page_with_options(&mut con, current_cursor, pattern, count, include_types).await?; - if !result.keys.is_empty() { - let next_cursor = if result.cursor != 0 { - encode_cluster_cursor(index, result.cursor)? - } else if index + 1 < master_nodes.len() { - encode_cluster_cursor(index + 1, 0)? - } else { - 0 + let mut connections: Vec> = std::iter::repeat_with(|| None).take(session.master_nodes.len()).collect(); + let mut total_keys = 0; + if include_total_keys { + for (index, endpoint) in session.master_nodes.iter().cloned().enumerate() { + let Ok(mut connection) = connect_node(endpoint).await else { + continue; }; - return Ok(RedisScanResult { cursor: next_cursor, keys: result.keys, total_keys }); - } - if result.cursor != 0 { - return Ok(RedisScanResult { - cursor: encode_cluster_cursor(index, result.cursor)?, - keys: Vec::new(), - total_keys, - }); + total_keys += redis::cmd("DBSIZE").query_async::(&mut connection).await.unwrap_or(0); + connections[index] = Some(connection); } } - Ok(RedisScanResult { cursor: 0, keys: Vec::new(), total_keys }) + let iterations = max_iterations.max(1); + let target_keys = count.max(1); + let mut all_keys = Vec::new(); + + for _ in 0..iterations { + if connections[session.node_index].is_none() { + connections[session.node_index] = + Some(connect_node(session.master_nodes[session.node_index].clone()).await?); + } + let connection = connections[session.node_index].as_mut().ok_or("Redis cluster node connection unavailable")?; + let page = + scan_keys_batch_inner(connection, session.node_cursor, pattern, count, 1, include_types, false).await?; + all_keys.extend(page.keys); + + let complete = if page.cursor != 0 { + session.node_cursor = page.cursor; + false + } else if session.node_index + 1 < session.master_nodes.len() { + session.node_index += 1; + session.node_cursor = 0; + false + } else { + true + }; + + if complete || all_keys.len() >= target_keys { + return Ok(RedisClusterKeyScanBatch { keys: all_keys, total_keys, complete }); + } + } + + Ok(RedisClusterKeyScanBatch { keys: all_keys, total_keys, complete: false }) +} + +fn canonical_cluster_master_nodes(mut master_nodes: Vec) -> Vec { + master_nodes.sort_unstable_by(|left, right| left.host.cmp(&right.host).then(left.port.cmp(&right.port))); + master_nodes.dedup(); + master_nodes } pub async fn scan_cluster_values_page( @@ -1783,11 +1982,27 @@ pub async fn scan_keys_batch( max_iterations: usize, include_types: bool, ) -> Result +where + C: ConnectionLike + Send + Sync + Unpin, +{ + scan_keys_batch_inner(con, cursor, pattern, count, max_iterations, include_types, true).await +} + +async fn scan_keys_batch_inner( + con: &mut C, + cursor: u64, + pattern: &str, + count: usize, + max_iterations: usize, + include_types: bool, + include_total_keys: bool, +) -> Result where C: ConnectionLike + Send + Sync + Unpin, { let iterations = max_iterations.max(1); - let total_keys: u64 = if cursor == 0 { redis::cmd("DBSIZE").query_async(con).await.unwrap_or(0) } else { 0 }; + let total_keys: u64 = + if include_total_keys && cursor == 0 { redis::cmd("DBSIZE").query_async(con).await.unwrap_or(0) } else { 0 }; let is_exact_match = !pattern.contains('*') && !pattern.contains('?') && !pattern.contains('['); if cursor == 0 && is_exact_match && !pattern.is_empty() { @@ -2718,7 +2933,14 @@ fn parse_scan_members(raw: RedisRawValue) -> Result<(u64, Vec), St #[cfg(test)] mod tests { - use std::collections::VecDeque; + use std::{ + collections::{HashMap, VecDeque}, + future::ready, + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, Mutex as StdMutex, + }, + }; use super::{ classify_command, connection_info, decode_cluster_cursor, encode_cluster_cursor, is_redis_json_type, @@ -2728,7 +2950,7 @@ mod tests { redis_key_matches_query, redis_key_raw_to_bytes, redis_key_value_preview, redis_sentinel_master_endpoint, redis_value_matches_query, redis_value_to_bytes, RedisAuthCandidate, RedisBlob, RedisBlobEncoding, RedisClusterSlotRange, RedisCollectionPage, RedisCommandSafety, RedisHashItem, RedisNodeEndpoint, - RedisRawValue, RedisSetItem, RedisStreamEntry, RedisStreamField, RedisValue, RedisValueData, + RedisNodeRoute, RedisRawValue, RedisSetItem, RedisStreamEntry, RedisStreamField, RedisValue, RedisValueData, }; use crate::models::connection::ConnectionConfig; use redis::{aio::ConnectionLike, Cmd, ConnectionAddr, Pipeline, RedisFuture}; @@ -2774,6 +2996,43 @@ mod tests { } } + struct TrackedRedisConnection { + responses: VecDeque>, + commands: Arc>>, + } + + impl TrackedRedisConnection { + fn new(responses: Vec, commands: Arc>>) -> Self { + Self { responses: responses.into_iter().map(Ok).collect(), commands } + } + } + + impl ConnectionLike for TrackedRedisConnection { + fn req_packed_command<'a>(&'a mut self, cmd: &'a Cmd) -> RedisFuture<'a, RedisRawValue> { + self.commands.lock().unwrap().push(String::from_utf8_lossy(&cmd.get_packed_command()).into_owned()); + let response = self.responses.pop_front().unwrap_or(Ok(RedisRawValue::Nil)); + Box::pin(async move { response }) + } + + fn req_packed_commands<'a>( + &'a mut self, + _cmd: &'a Pipeline, + _offset: usize, + _count: usize, + ) -> RedisFuture<'a, Vec> { + Box::pin(async move { Ok(Vec::new()) }) + } + + fn get_db(&self) -> i64 { + 0 + } + } + + fn tracked_command_count(commands: &Arc>>, command: &str) -> usize { + let needle = format!("\r\n{command}\r\n"); + commands.lock().unwrap().iter().filter(|packed| packed.contains(&needle)).count() + } + fn bulk(value: &str) -> RedisRawValue { RedisRawValue::BulkString(value.as_bytes().to_vec()) } @@ -3038,6 +3297,454 @@ mod tests { assert_eq!(con.command_count("SCAN"), 2); } + #[tokio::test] + async fn cluster_key_batch_reuses_node_connections_across_scan_iterations() { + let sessions = tokio::sync::Mutex::new(super::RedisClusterScanSessions::default()); + let masters = vec![ + RedisNodeEndpoint { host: "node-a".to_string(), port: 7000 }, + RedisNodeEndpoint { host: "node-b".to_string(), port: 7001 }, + RedisNodeEndpoint { host: "node-c".to_string(), port: 7002 }, + ]; + let logs: Vec<_> = (0..masters.len()).map(|_| Arc::new(StdMutex::new(Vec::new()))).collect(); + let mut connections = HashMap::from([ + ( + masters[0].clone(), + TrackedRedisConnection::new( + vec![RedisRawValue::Int(100), scan_response("7", vec![]), scan_response("0", vec![])], + logs[0].clone(), + ), + ), + ( + masters[1].clone(), + TrackedRedisConnection::new( + vec![ + RedisRawValue::Int(200), + scan_response("9", vec![]), + scan_response("0", vec!["membership:saas:base:token:match"]), + ], + logs[1].clone(), + ), + ), + (masters[2].clone(), TrackedRedisConnection::new(vec![RedisRawValue::Int(300)], logs[2].clone())), + ]); + let connect_count = Arc::new(AtomicUsize::new(0)); + let connector_count = connect_count.clone(); + let topology_count = Arc::new(AtomicUsize::new(0)); + let discovery_count = topology_count.clone(); + let discovered_masters = masters.clone(); + + let result = super::scan_cluster_keys_batch_with( + &sessions, + move || { + discovery_count.fetch_add(1, Ordering::Relaxed); + ready(Ok(discovered_masters.clone())) + }, + move |endpoint| { + connector_count.fetch_add(1, Ordering::Relaxed); + ready(connections.remove(&endpoint).ok_or_else(|| format!("unexpected reconnect to {}", endpoint.host))) + }, + 0, + "membership:saas:base:token:*", + 1000, + 4, + false, + ) + .await + .unwrap(); + + assert_eq!(result.total_keys, 600); + assert_eq!(result.keys.len(), 1); + assert!(result.cursor > 0); + assert!(result.cursor <= super::MAX_SAFE_INTEGER_CURSOR); + assert_eq!(topology_count.load(Ordering::Relaxed), 1); + assert_eq!(connect_count.load(Ordering::Relaxed), 3); + assert_eq!(logs.iter().map(|log| tracked_command_count(log, "DBSIZE")).sum::(), 3); + assert_eq!(logs.iter().map(|log| tracked_command_count(log, "SCAN")).sum::(), 4); + assert_eq!(sessions.lock().await.entries.len(), 1); + } + + #[tokio::test] + async fn cluster_key_batch_continuation_skips_dbsize_and_advances_masters() { + let pattern = "membership:saas:base:token:*"; + let masters = vec![ + RedisNodeEndpoint { host: "node-a".to_string(), port: 7000 }, + RedisNodeEndpoint { host: "node-b".to_string(), port: 7001 }, + RedisNodeEndpoint { host: "node-c".to_string(), port: 7002 }, + ]; + let mut saved_session = super::RedisClusterKeyScanSession::new(masters.clone(), pattern); + saved_session.node_index = 1; + saved_session.node_cursor = 7; + let mut saved_sessions = super::RedisClusterScanSessions::default(); + let cursor = saved_sessions.insert(None, saved_session); + let sessions = tokio::sync::Mutex::new(saved_sessions); + let logs: Vec<_> = (0..masters.len()).map(|_| Arc::new(StdMutex::new(Vec::new()))).collect(); + let mut connections = HashMap::from([ + (masters[1].clone(), TrackedRedisConnection::new(vec![scan_response("0", vec![])], logs[1].clone())), + ( + masters[2].clone(), + TrackedRedisConnection::new( + vec![scan_response("0", vec!["membership:saas:base:token:last"])], + logs[2].clone(), + ), + ), + ]); + let connect_count = Arc::new(AtomicUsize::new(0)); + let connector_count = connect_count.clone(); + let topology_count = Arc::new(AtomicUsize::new(0)); + let discovery_count = topology_count.clone(); + + let result = super::scan_cluster_keys_batch_with( + &sessions, + move || { + discovery_count.fetch_add(1, Ordering::Relaxed); + ready(Err("continuation must not rediscover topology".to_string())) + }, + move |endpoint| { + connector_count.fetch_add(1, Ordering::Relaxed); + ready( + connections.remove(&endpoint).ok_or_else(|| format!("unexpected connection to {}", endpoint.host)), + ) + }, + cursor, + pattern, + 1000, + 2, + false, + ) + .await + .unwrap(); + + assert_eq!(result.cursor, 0); + assert_eq!(result.total_keys, 0); + assert_eq!(result.keys.len(), 1); + assert_eq!(topology_count.load(Ordering::Relaxed), 0); + assert_eq!(connect_count.load(Ordering::Relaxed), 2); + assert_eq!(logs.iter().map(|log| tracked_command_count(log, "DBSIZE")).sum::(), 0); + assert_eq!(logs.iter().map(|log| tracked_command_count(log, "SCAN")).sum::(), 2); + assert!(sessions.lock().await.entries.is_empty()); + } + + #[tokio::test] + async fn cluster_key_batch_refreshes_healthy_cached_masters_before_scanning() { + let pattern = "membership:saas:base:token:*"; + let cached_masters = [ + RedisNodeEndpoint { host: "node-a".to_string(), port: 7000 }, + RedisNodeEndpoint { host: "node-b".to_string(), port: 7001 }, + ]; + let current_masters = [ + cached_masters[0].clone(), + cached_masters[1].clone(), + RedisNodeEndpoint { host: "node-c".to_string(), port: 7002 }, + ]; + let sessions = tokio::sync::Mutex::new(super::RedisClusterScanSessions::default()); + let logs: Vec<_> = (0..current_masters.len()).map(|_| Arc::new(StdMutex::new(Vec::new()))).collect(); + let mut connections = HashMap::from([ + ( + current_masters[0].clone(), + TrackedRedisConnection::new(vec![RedisRawValue::Int(100), scan_response("0", vec![])], logs[0].clone()), + ), + ( + current_masters[1].clone(), + TrackedRedisConnection::new(vec![RedisRawValue::Int(200), scan_response("0", vec![])], logs[1].clone()), + ), + ( + current_masters[2].clone(), + TrackedRedisConnection::new( + vec![RedisRawValue::Int(300), scan_response("0", vec!["membership:saas:base:token:new-master"])], + logs[2].clone(), + ), + ), + ]); + let discovered_masters = vec![ + current_masters[2].clone(), + current_masters[0].clone(), + current_masters[1].clone(), + current_masters[2].clone(), + ]; + let topology_count = Arc::new(AtomicUsize::new(0)); + let discovery_count = topology_count.clone(); + + let result = super::scan_cluster_keys_batch_with( + &sessions, + move || { + discovery_count.fetch_add(1, Ordering::Relaxed); + ready(Ok(discovered_masters.clone())) + }, + move |endpoint| { + ready(connections.remove(&endpoint).ok_or_else(|| format!("unexpected reconnect to {}", endpoint.host))) + }, + 0, + pattern, + 1000, + 3, + false, + ) + .await + .unwrap(); + + assert_eq!(result.cursor, 0); + assert_eq!(result.total_keys, 600); + assert_eq!(result.keys.len(), 1); + assert_eq!(result.keys[0].key_display, "membership:saas:base:token:new-master"); + assert_eq!(topology_count.load(Ordering::Relaxed), 1); + assert_eq!(logs.iter().map(|log| tracked_command_count(log, "DBSIZE")).sum::(), 3); + assert_eq!(logs.iter().map(|log| tracked_command_count(log, "SCAN")).sum::(), 3); + assert!(sessions.lock().await.entries.is_empty()); + } + + #[test] + fn cluster_scan_routes_discovered_advertised_endpoint_through_existing_mapping() { + let advertised = RedisNodeEndpoint { host: "redis.internal".to_string(), port: 7000 }; + let forwarded = RedisNodeEndpoint { host: "127.0.0.1".to_string(), port: 17_000 }; + let pool = super::RedisClusterPool { + connection: None, + seed_nodes: vec![advertised.clone()], + seed_routes: Vec::new(), + slot_ranges: Vec::new(), + node_routes: vec![RedisNodeRoute { advertised: advertised.clone(), connect: forwarded.clone() }], + tls: false, + tls_insecure: false, + username: String::new(), + password: String::new(), + scan_sessions: Box::new(tokio::sync::Mutex::new(super::RedisClusterScanSessions::default())), + }; + + assert_eq!(super::mapped_cluster_endpoint(&pool, &advertised), forwarded); + } + + #[tokio::test] + async fn cluster_key_batch_continuation_uses_stable_node_identity_without_rediscovery() { + let pattern = "membership:*"; + let masters = vec![ + RedisNodeEndpoint { host: "node-a".to_string(), port: 7000 }, + RedisNodeEndpoint { host: "node-b".to_string(), port: 7001 }, + RedisNodeEndpoint { host: "node-c".to_string(), port: 7002 }, + ]; + let mut saved_session = super::RedisClusterKeyScanSession::new(masters.clone(), pattern); + saved_session.node_index = 1; + saved_session.node_cursor = 7; + let mut saved_sessions = super::RedisClusterScanSessions::default(); + let cursor = saved_sessions.insert(None, saved_session); + let sessions = tokio::sync::Mutex::new(saved_sessions); + let commands = Arc::new(StdMutex::new(Vec::new())); + let mut connections = HashMap::from([( + masters[1].clone(), + TrackedRedisConnection::new(vec![scan_response("9", vec!["membership:stable-node"])], commands.clone()), + )]); + + let result = super::scan_cluster_keys_batch_with( + &sessions, + || ready(Err("continuation must not rediscover topology".to_string())), + move |endpoint| { + ready( + connections.remove(&endpoint).ok_or_else(|| format!("unexpected connection to {}", endpoint.host)), + ) + }, + cursor, + pattern, + 1, + 1, + false, + ) + .await + .unwrap(); + + assert_eq!(result.cursor, cursor); + assert_eq!(result.total_keys, 0); + assert_eq!(result.keys[0].key_display, "membership:stable-node"); + assert_eq!(tracked_command_count(&commands, "DBSIZE"), 0); + assert_eq!(tracked_command_count(&commands, "SCAN"), 1); + let continued = sessions.lock().await.take(cursor).unwrap(); + assert_eq!(continued.master_nodes[continued.node_index], masters[1]); + assert_eq!(continued.node_cursor, 9); + } + + #[tokio::test] + async fn cluster_key_batch_continuation_restores_original_position_after_failure() { + let pattern = "membership:*"; + let masters = vec![ + RedisNodeEndpoint { host: "node-a".to_string(), port: 7000 }, + RedisNodeEndpoint { host: "node-b".to_string(), port: 7001 }, + RedisNodeEndpoint { host: "node-c".to_string(), port: 7002 }, + ]; + let mut saved_session = super::RedisClusterKeyScanSession::new(masters.clone(), pattern); + saved_session.node_index = 1; + saved_session.node_cursor = 7; + let mut saved_sessions = super::RedisClusterScanSessions::default(); + let cursor = saved_sessions.insert(None, saved_session); + let sessions = tokio::sync::Mutex::new(saved_sessions); + let commands = Arc::new(StdMutex::new(Vec::new())); + let connected_nodes = Arc::new(StdMutex::new(Vec::new())); + let first_connected_nodes = connected_nodes.clone(); + let first_commands = commands.clone(); + let first_master = masters[1].clone(); + let failed_master = masters[2].clone(); + let mut first_connection = Some(TrackedRedisConnection::new( + vec![scan_response("0", vec!["membership:before-failure"])], + first_commands, + )); + + let error = super::scan_cluster_keys_batch_with( + &sessions, + || ready(Err("continuation must not rediscover topology".to_string())), + move |endpoint| { + first_connected_nodes.lock().unwrap().push(endpoint.clone()); + ready(if endpoint == first_master { + first_connection.take().ok_or_else(|| "unexpected reconnect".to_string()) + } else if endpoint == failed_master { + Err("temporary node failure".to_string()) + } else { + Err(format!("unexpected connection to {}", endpoint.host)) + }) + }, + cursor, + pattern, + 1000, + 2, + false, + ) + .await + .unwrap_err(); + + assert_eq!(error, "temporary node failure"); + let restored = sessions.lock().await.entries.get(&cursor).unwrap().clone(); + assert_eq!(restored.master_nodes[restored.node_index], masters[1]); + assert_eq!(restored.node_cursor, 7); + + let retry_connected_nodes = connected_nodes.clone(); + let retry_commands = commands.clone(); + let retry_master = masters[1].clone(); + let mut retry_connection = + Some(TrackedRedisConnection::new(vec![scan_response("9", vec!["membership:retried"])], retry_commands)); + let result = super::scan_cluster_keys_batch_with( + &sessions, + || ready(Err("continuation must not rediscover topology".to_string())), + move |endpoint| { + retry_connected_nodes.lock().unwrap().push(endpoint.clone()); + ready(if endpoint == retry_master { + retry_connection.take().ok_or_else(|| "unexpected reconnect".to_string()) + } else { + Err(format!("unexpected connection to {}", endpoint.host)) + }) + }, + cursor, + pattern, + 1, + 1, + false, + ) + .await + .unwrap(); + + assert_eq!(result.cursor, cursor); + assert_eq!(result.total_keys, 0); + assert_eq!(result.keys[0].key_display, "membership:retried"); + assert_eq!(tracked_command_count(&commands, "DBSIZE"), 0); + assert_eq!(tracked_command_count(&commands, "SCAN"), 2); + assert!(commands.lock().unwrap()[1].contains("\r\n7\r\n")); + assert_eq!(*connected_nodes.lock().unwrap(), vec![masters[1].clone(), masters[2].clone(), masters[1].clone()]); + let continued = sessions.lock().await.take(cursor).unwrap(); + assert_eq!(continued.master_nodes[continued.node_index], masters[1]); + assert_eq!(continued.node_cursor, 9); + } + + #[tokio::test] + async fn cluster_key_batch_rejects_an_unknown_continuation_cursor() { + let sessions = tokio::sync::Mutex::new(super::RedisClusterScanSessions::default()); + let topology_count = Arc::new(AtomicUsize::new(0)); + let discovery_count = topology_count.clone(); + let connect_count = Arc::new(AtomicUsize::new(0)); + let connector_count = connect_count.clone(); + + let error = super::scan_cluster_keys_batch_with::( + &sessions, + move || { + discovery_count.fetch_add(1, Ordering::Relaxed); + ready(Err("continuation must not rediscover topology".to_string())) + }, + move |_| { + connector_count.fetch_add(1, Ordering::Relaxed); + ready(Err("continuation must not connect".to_string())) + }, + 999, + "membership:*", + 1000, + 1, + false, + ) + .await + .unwrap_err(); + + assert_eq!(error, super::INVALID_CLUSTER_SCAN_CURSOR_ERROR); + assert_eq!(topology_count.load(Ordering::Relaxed), 0); + assert_eq!(connect_count.load(Ordering::Relaxed), 0); + } + + #[tokio::test] + async fn cluster_key_batch_rejects_an_expired_continuation_cursor() { + let master = RedisNodeEndpoint { host: "node-a".to_string(), port: 7000 }; + let mut saved_sessions = super::RedisClusterScanSessions::default(); + let cursor = saved_sessions.insert(None, super::RedisClusterKeyScanSession::new(vec![master], "membership:*")); + saved_sessions.entries.get_mut(&cursor).unwrap().last_used = + std::time::Instant::now() - super::CLUSTER_SCAN_SESSION_TTL - std::time::Duration::from_secs(1); + let sessions = tokio::sync::Mutex::new(saved_sessions); + + let error = super::scan_cluster_keys_batch_with::( + &sessions, + || ready(Err("continuation must not rediscover topology".to_string())), + |_| ready(Err("continuation must not connect".to_string())), + cursor, + "membership:*", + 1000, + 1, + false, + ) + .await + .unwrap_err(); + + assert_eq!(error, super::INVALID_CLUSTER_SCAN_CURSOR_ERROR); + assert!(sessions.lock().await.entries.is_empty()); + } + + #[tokio::test] + async fn cluster_key_batch_reports_new_scan_topology_failure() { + let sessions = tokio::sync::Mutex::new(super::RedisClusterScanSessions::default()); + + let error = super::scan_cluster_keys_batch_with::( + &sessions, + || ready(Err("topology unavailable".to_string())), + |_| ready(Err("scan connection must not be opened".to_string())), + 0, + "membership:*", + 1000, + 1, + false, + ) + .await + .unwrap_err(); + + assert_eq!(error, "topology unavailable"); + assert!(sessions.lock().await.entries.is_empty()); + } + + #[test] + fn cluster_scan_sessions_expire_and_remain_bounded() { + let master = RedisNodeEndpoint { host: "node-a".to_string(), port: 7000 }; + let mut sessions = super::RedisClusterScanSessions::default(); + let expired_cursor = sessions.insert(None, super::RedisClusterKeyScanSession::new(vec![master.clone()], "*")); + sessions.entries.get_mut(&expired_cursor).unwrap().last_used = + std::time::Instant::now() - super::CLUSTER_SCAN_SESSION_TTL - std::time::Duration::from_secs(1); + + for index in 0..=super::CLUSTER_SCAN_SESSION_LIMIT { + sessions + .insert(None, super::RedisClusterKeyScanSession::new(vec![master.clone()], &format!("key:{index}:*"))); + } + + assert!(!sessions.entries.contains_key(&expired_cursor)); + assert_eq!(sessions.entries.len(), super::CLUSTER_SCAN_SESSION_LIMIT); + assert!(sessions.entries.keys().all(|cursor| *cursor > 0 && *cursor <= super::MAX_SAFE_INTEGER_CURSOR)); + } + #[tokio::test] async fn filtered_hash_load_more_matches_fields_and_keeps_scan_cursor() { let mut con = FakeRedisConnection::new(vec![ diff --git a/crates/dbx-core/src/redis_ops.rs b/crates/dbx-core/src/redis_ops.rs index e5da5ac36..b257c253b 100644 --- a/crates/dbx-core/src/redis_ops.rs +++ b/crates/dbx-core/src/redis_ops.rs @@ -1,7 +1,6 @@ use crate::connection::{AppState, PoolKind}; use crate::db::redis_driver::{ - self, RedisCollectionPage, RedisCommandResult, RedisConnection, RedisDatabaseInfo, RedisKeyInfo, RedisScanResult, - RedisValue, + self, RedisCollectionPage, RedisCommandResult, RedisConnection, RedisDatabaseInfo, RedisScanResult, RedisValue, }; async fn ensure_redis_pool(state: &AppState, connection_id: &str) -> Result<(), String> { @@ -65,40 +64,8 @@ pub async fn redis_scan_keys_batch_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - // Cluster scan already iterates across nodes; for batch mode we - // loop the cluster-level scan to accumulate keys server-side. - if max_iterations <= 1 { - return redis_driver::scan_cluster_keys_page_with_options( - cluster, - cursor, - pattern, - count, - include_types, - ) - .await; - } - let mut all_keys: Vec = Vec::new(); - let mut current_cursor = cursor; - let mut total_keys: u64 = 0; - for i in 0..max_iterations { - let page = redis_driver::scan_cluster_keys_page_with_options( - cluster, - current_cursor, - pattern, - count, - include_types, - ) - .await?; - if i == 0 { - total_keys = page.total_keys; - } - all_keys.extend(page.keys); - if page.cursor == 0 { - return Ok(RedisScanResult { cursor: 0, keys: all_keys, total_keys }); - } - current_cursor = page.cursor; - } - Ok(RedisScanResult { cursor: current_cursor, keys: all_keys, total_keys }) + redis_driver::scan_cluster_keys_batch(cluster, cursor, pattern, count, max_iterations, include_types) + .await } }, _ => Err("Not a Redis connection".to_string()),