perf(redis): reuse cluster connections during key scans

This commit is contained in:
Guoyu Su 2026-07-21 21:36:06 +08:00 committed by GitHub
parent 81e23ab870
commit 9118240cd4
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 741 additions and 67 deletions

View File

@ -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<Mutex<RedisClusterScanSessions>>,
}
#[derive(Debug, Clone)]
struct RedisClusterKeyScanSession {
master_nodes: Vec<RedisNodeEndpoint>,
node_index: usize,
node_cursor: u64,
pattern: String,
last_used: Instant,
}
impl RedisClusterKeyScanSession {
fn new(master_nodes: Vec<RedisNodeEndpoint>, 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<u64, RedisClusterKeyScanSession>,
}
impl Default for RedisClusterScanSessions {
fn default() -> Self {
Self { next_cursor: 1, entries: HashMap::new() }
}
}
impl RedisClusterScanSessions {
fn take(&mut self, cursor: u64) -> Option<RedisClusterKeyScanSession> {
self.remove_expired();
self.entries.remove(&cursor)
}
fn insert(&mut self, cursor: Option<u64>, 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<RedisClusterPo
tls_insecure: config.redis_tls_insecure(),
username: auth.username,
password: auth.password,
scan_sessions: Box::new(Mutex::new(RedisClusterScanSessions::default())),
});
}
Err(err) if last_error.is_none() || is_redis_auth_error(&err) => {
@ -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<RedisScanResult, String> {
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<RedisScanResult, String> {
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<C, Discover, DiscoverFuture, Connect, ConnectFuture>(
sessions: &Mutex<RedisClusterScanSessions>,
mut discover_master_nodes: Discover,
connect_node: Connect,
cursor: u64,
pattern: &str,
count: usize,
max_iterations: usize,
include_types: bool,
) -> Result<RedisScanResult, String>
where
C: ConnectionLike + Send + Sync + Unpin,
Discover: FnMut() -> DiscoverFuture,
DiscoverFuture: Future<Output = Result<Vec<RedisNodeEndpoint>, String>>,
Connect: FnMut(RedisNodeEndpoint) -> ConnectFuture,
ConnectFuture: Future<Output = Result<C, String>>,
{
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<RedisKeyInfo>,
total_keys: u64,
complete: bool,
}
async fn scan_cluster_keys_batch_on_session<C, Connect, ConnectFuture>(
session: &mut RedisClusterKeyScanSession,
mut connect_node: Connect,
pattern: &str,
count: usize,
max_iterations: usize,
include_types: bool,
include_total_keys: bool,
) -> Result<RedisClusterKeyScanBatch, String>
where
C: ConnectionLike + Send + Sync + Unpin,
Connect: FnMut(RedisNodeEndpoint) -> ConnectFuture,
ConnectFuture: Future<Output = Result<C, String>>,
{
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<Option<C>> = 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::<u64>(&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<RedisNodeEndpoint>) -> Vec<RedisNodeEndpoint> {
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<C>(
max_iterations: usize,
include_types: bool,
) -> Result<RedisScanResult, String>
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<C>(
con: &mut C,
cursor: u64,
pattern: &str,
count: usize,
max_iterations: usize,
include_types: bool,
include_total_keys: bool,
) -> Result<RedisScanResult, String>
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<RedisSetItem>), 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<redis::RedisResult<RedisRawValue>>,
commands: Arc<StdMutex<Vec<String>>>,
}
impl TrackedRedisConnection {
fn new(responses: Vec<RedisRawValue>, commands: Arc<StdMutex<Vec<String>>>) -> 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<RedisRawValue>> {
Box::pin(async move { Ok(Vec::new()) })
}
fn get_db(&self) -> i64 {
0
}
}
fn tracked_command_count(commands: &Arc<StdMutex<Vec<String>>>, 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::<usize>(), 3);
assert_eq!(logs.iter().map(|log| tracked_command_count(log, "SCAN")).sum::<usize>(), 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::<usize>(), 0);
assert_eq!(logs.iter().map(|log| tracked_command_count(log, "SCAN")).sum::<usize>(), 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::<usize>(), 3);
assert_eq!(logs.iter().map(|log| tracked_command_count(log, "SCAN")).sum::<usize>(), 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::<TrackedRedisConnection, _, _, _, _>(
&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::<TrackedRedisConnection, _, _, _, _>(
&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::<TrackedRedisConnection, _, _, _, _>(
&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![

View File

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