diff --git a/crates/dbx-core/src/connection.rs b/crates/dbx-core/src/connection.rs index 72639a3be..209147de4 100644 --- a/crates/dbx-core/src/connection.rs +++ b/crates/dbx-core/src/connection.rs @@ -489,7 +489,9 @@ impl AppState { } DatabaseType::Redis => { let con = if db_config.uses_redis_cluster() { - db::redis_driver::RedisConnection::Cluster(db::redis_driver::connect_cluster(&db_config).await?) + db::redis_driver::RedisConnection::Cluster( + self.connect_redis_cluster(connection_id, &db_config).await?, + ) } else if db_config.uses_redis_sentinel() { db::redis_driver::RedisConnection::Direct(tokio::sync::Mutex::new( db::redis_driver::connect_sentinel(&db_config).await?, @@ -801,6 +803,63 @@ impl AppState { Ok(("127.0.0.1".to_string(), local_port)) } + pub async fn connect_redis_cluster( + &self, + connection_id: &str, + config: &ConnectionConfig, + ) -> Result { + let transport_layers = config.effective_transport_layers(); + if transport_layers.is_empty() { + return db::redis_driver::connect_cluster(config).await; + } + + let result = async { + let seed_nodes = db::redis_driver::redis_cluster_seed_nodes(config)?; + let seed_routes = self.redis_cluster_node_routes(connection_id, &transport_layers, &seed_nodes).await?; + let (auth, slot_ranges) = + db::redis_driver::discover_cluster_slot_ranges_from_routes(config, &seed_routes).await?; + let master_nodes = db::redis_driver::unique_master_nodes(&slot_ranges); + let node_routes = self.redis_cluster_node_routes(connection_id, &transport_layers, &master_nodes).await?; + + db::redis_driver::connect_routed_cluster(config, seed_routes, slot_ranges, node_routes, auth).await + } + .await; + + if result.is_err() { + let redis_cluster_prefix = redis_cluster_transport_prefix(connection_id); + self.tunnels.stop_tunnels_with_prefix(&redis_cluster_prefix).await; + self.proxy_tunnels.stop_tunnels_with_prefix(&redis_cluster_prefix).await; + } + + result + } + + async fn redis_cluster_node_routes( + &self, + connection_id: &str, + transport_layers: &[crate::models::connection::TransportLayerConfig], + nodes: &[db::redis_driver::RedisNodeEndpoint], + ) -> Result, String> { + let mut routes = Vec::with_capacity(nodes.len()); + for node in nodes { + let tunnel_id = redis_cluster_transport_id(connection_id, node); + let local_port = db::transport_layer_tunnel::start_transport_layers( + &tunnel_id, + transport_layers, + &node.host, + node.port, + &self.tunnels, + &self.proxy_tunnels, + ) + .await?; + routes.push(db::redis_driver::RedisNodeRoute { + advertised: node.clone(), + connect: db::redis_driver::RedisNodeEndpoint { host: "127.0.0.1".to_string(), port: local_port }, + }); + } + Ok(routes) + } + #[cfg(feature = "mq-admin")] pub async fn mq_admin_config_for_connection( &self, @@ -1076,6 +1135,9 @@ impl AppState { } async fn reset_connection_transport_layers(&self, connection_id: &str, layer_count: usize) { + let redis_cluster_prefix = redis_cluster_transport_prefix(connection_id); + self.tunnels.stop_tunnels_with_prefix(&redis_cluster_prefix).await; + self.proxy_tunnels.stop_tunnels_with_prefix(&redis_cluster_prefix).await; db::transport_layer_tunnel::stop_transport_layers( connection_id, layer_count, @@ -1290,6 +1352,19 @@ fn normalize_client_session_id(client_session_id: Option<&str>) -> Option String { + format!("{connection_id}:redis-cluster:") +} + +fn redis_cluster_transport_id(connection_id: &str, endpoint: &db::redis_driver::RedisNodeEndpoint) -> String { + format!( + "{}{host}:{port}", + redis_cluster_transport_prefix(connection_id), + host = endpoint.host, + port = endpoint.port + ) +} + fn session_scoped_pool_key(base_pool_key: String, client_session_id: Option<&str>) -> String { normalize_client_session_id(client_session_id) .map(|session| format!("{base_pool_key}:session:{session}")) diff --git a/crates/dbx-core/src/db/proxy_tunnel.rs b/crates/dbx-core/src/db/proxy_tunnel.rs index 2652402d5..5a927329b 100644 --- a/crates/dbx-core/src/db/proxy_tunnel.rs +++ b/crates/dbx-core/src/db/proxy_tunnel.rs @@ -69,6 +69,16 @@ impl ProxyTunnelManager { handle.abort(); } } + + pub async fn stop_tunnels_with_prefix(&self, connection_id_prefix: &str) { + let mut tunnels = self.tunnels.lock().await; + let keys: Vec = tunnels.keys().filter(|key| key.starts_with(connection_id_prefix)).cloned().collect(); + for key in keys { + if let Some((handle, _)) = tunnels.remove(&key) { + handle.abort(); + } + } + } } #[derive(Clone)] diff --git a/crates/dbx-core/src/db/redis_driver.rs b/crates/dbx-core/src/db/redis_driver.rs index de68f5833..38cff78f8 100644 --- a/crates/dbx-core/src/db/redis_driver.rs +++ b/crates/dbx-core/src/db/redis_driver.rs @@ -5,11 +5,11 @@ use redis::{ cluster::ClusterClient, cluster_async::ClusterConnection, sentinel::{Sentinel, SentinelNodeConnectionInfo}, - ConnectionAddr, ConnectionInfo, FromRedisValue, ProtocolVersion, RedisConnectionInfo, TlsMode, - Value as RedisRawValue, + Cmd, ConnectionAddr, ConnectionInfo, FromRedisValue, Pipeline, ProtocolVersion, RedisConnectionInfo, RedisFuture, + TlsMode, Value as RedisRawValue, }; use serde::{Deserialize, Serialize}; -use tokio::sync::Mutex; +use tokio::sync::{Mutex, MutexGuard}; const STREAM_ENTRY_LIMIT: usize = 100; const COLLECTION_PAGE_SIZE: usize = 200; @@ -81,14 +81,41 @@ pub enum RedisConnection { } pub struct RedisClusterPool { - pub connection: Mutex, + pub connection: Option>, pub seed_nodes: Vec, + pub seed_routes: Vec, + pub slot_ranges: Vec, + pub node_routes: Vec, pub tls: bool, pub tls_insecure: bool, pub username: String, pub password: String, } +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RedisClusterAuth { + pub username: String, + pub password: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RedisNodeRoute { + pub advertised: RedisNodeEndpoint, + pub connect: RedisNodeEndpoint, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RedisClusterSlotRange { + pub start: u16, + pub end: u16, + pub master: RedisNodeEndpoint, +} + +pub enum RedisClusterConnectionGuard<'a> { + Native(MutexGuard<'a, ClusterConnection>), + Direct(redis::aio::MultiplexedConnection), +} + #[derive(Debug, Clone, PartialEq, Eq)] struct RedisAuthCandidate { username: String, @@ -101,6 +128,34 @@ pub struct RedisNodeEndpoint { pub port: u16, } +impl ConnectionLike for RedisClusterConnectionGuard<'_> { + fn req_packed_command<'a>(&'a mut self, cmd: &'a Cmd) -> RedisFuture<'a, RedisRawValue> { + match self { + RedisClusterConnectionGuard::Native(con) => con.req_packed_command(cmd), + RedisClusterConnectionGuard::Direct(con) => con.req_packed_command(cmd), + } + } + + fn req_packed_commands<'a>( + &'a mut self, + cmd: &'a Pipeline, + offset: usize, + count: usize, + ) -> RedisFuture<'a, Vec> { + match self { + RedisClusterConnectionGuard::Native(con) => con.req_packed_commands(cmd, offset, count), + RedisClusterConnectionGuard::Direct(con) => con.req_packed_commands(cmd, offset, count), + } + } + + fn get_db(&self) -> i64 { + match self { + RedisClusterConnectionGuard::Native(con) => con.get_db(), + RedisClusterConnectionGuard::Direct(con) => con.get_db(), + } + } +} + pub async fn connect(url: &str, timeout: std::time::Duration) -> Result { let client = redis::Client::open(url).map_err(|e| format!("Redis connection failed: {e}"))?; connect_client_with_timeout(client, timeout, "Redis").await @@ -222,9 +277,23 @@ pub async fn connect_cluster(config: &ConnectionConfig) -> Result { + let seed_routes = identity_routes(&seed_nodes); + let slot_ranges = cluster_slot_ranges_from_routes( + &seed_routes, + config.ssl, + config.redis_tls_insecure(), + &auth.username, + &auth.password, + ) + .await + .unwrap_or_default(); + let node_routes = identity_routes(&unique_master_nodes(&slot_ranges)); return Ok(RedisClusterPool { - connection: Mutex::new(con), + connection: Some(Mutex::new(con)), seed_nodes, + seed_routes, + slot_ranges, + node_routes, tls: config.ssl, tls_insecure: config.redis_tls_insecure(), username: auth.username, @@ -244,6 +313,69 @@ pub async fn connect_cluster(config: &ConnectionConfig) -> Result Result<(RedisClusterAuth, Vec), String> { + let mut last_error = None; + for auth in redis_auth_candidates(&config.username, &config.password) { + match cluster_slot_ranges_from_routes( + seed_routes, + config.ssl, + config.redis_tls_insecure(), + &auth.username, + &auth.password, + ) + .await + { + Ok(slot_ranges) if !slot_ranges.is_empty() => { + return Ok((RedisClusterAuth { username: auth.username, password: auth.password }, slot_ranges)); + } + Ok(_) => { + last_error = Some("Redis cluster master discovery returned no slots".to_string()); + } + Err(err) if last_error.is_none() || is_redis_auth_error(&err) => { + let should_retry = is_redis_auth_error(&err); + last_error = Some(err); + if !should_retry { + break; + } + } + Err(err) => return Err(err), + } + } + Err(last_error.unwrap_or_else(|| "Redis cluster master discovery failed".to_string())) +} + +pub async fn connect_routed_cluster( + config: &ConnectionConfig, + seed_routes: Vec, + slot_ranges: Vec, + node_routes: Vec, + auth: RedisClusterAuth, +) -> Result { + let seed_nodes = redis_cluster_seed_nodes(config)?; + let mut pool = RedisClusterPool { + connection: None, + seed_nodes, + seed_routes, + slot_ranges, + node_routes, + tls: config.ssl, + tls_insecure: config.redis_tls_insecure(), + username: auth.username, + password: auth.password, + }; + if pool.node_routes.is_empty() { + pool.node_routes = identity_routes(&unique_master_nodes(&pool.slot_ranges)); + } + { + let mut con = cluster_any_connection(&pool).await?; + redis_ping(&mut con, "Redis cluster").await?; + } + Ok(pool) +} + pub async fn test_connection(connection: &RedisConnection) -> Result<(), String> { match connection { RedisConnection::Direct(con) => { @@ -251,8 +383,8 @@ pub async fn test_connection(connection: &RedisConnection) -> Result<(), String> redis_ping(&mut *con, "Redis").await } RedisConnection::Cluster(cluster) => { - let mut con = cluster.connection.lock().await; - redis_ping(&mut *con, "Redis cluster").await + let mut con = cluster_any_connection(cluster).await?; + redis_ping(&mut con, "Redis cluster").await } } } @@ -288,7 +420,7 @@ fn redis_sentinel_nodes(config: &ConnectionConfig) -> Result endpoints.iter().map(|endpoint| redis_sentinel_node_connection_info(config, endpoint)).collect() } -fn redis_cluster_seed_nodes(config: &ConnectionConfig) -> Result, String> { +pub fn redis_cluster_seed_nodes(config: &ConnectionConfig) -> Result, String> { redis_node_endpoints( config.redis_cluster_nodes.trim(), config.host.trim(), @@ -630,8 +762,7 @@ pub async fn scan_cluster_keys_page( 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_direct_node(endpoint, pool.tls, pool.tls_insecure, &pool.username, &pool.password).await?; + 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(&mut con, current_cursor, pattern, count).await?; if !result.keys.is_empty() { @@ -677,8 +808,7 @@ pub async fn scan_cluster_values_page( 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_direct_node(endpoint, pool.tls, pool.tls_insecure, &pool.username, &pool.password).await?; + let mut con = connect_cluster_node(pool, endpoint).await?; let current_cursor = if index == node_index { node_cursor } else { 0 }; let result = scan_values_page(&mut con, current_cursor, pattern, query, include_key_matches, count).await?; if !result.keys.is_empty() { @@ -704,14 +834,26 @@ pub async fn scan_cluster_values_page( } pub async fn cluster_master_nodes(pool: &RedisClusterPool) -> Result, String> { - cluster_master_nodes_from_seeds(&pool.seed_nodes, pool.tls, pool.tls_insecure, &pool.username, &pool.password).await + if !pool.seed_routes.is_empty() { + return cluster_slot_ranges_from_routes( + &pool.seed_routes, + pool.tls, + pool.tls_insecure, + &pool.username, + &pool.password, + ) + .await + .map(|slot_ranges| unique_master_nodes(&slot_ranges)); + } + cluster_slot_ranges_from_seeds(&pool.seed_nodes, pool.tls, pool.tls_insecure, &pool.username, &pool.password) + .await + .map(|slot_ranges| unique_master_nodes(&slot_ranges)) } pub async fn flush_cluster(pool: &RedisClusterPool) -> Result<(), String> { let master_nodes = cluster_master_nodes(pool).await?; for endpoint in master_nodes { - let mut con = - connect_direct_node(&endpoint, pool.tls, pool.tls_insecure, &pool.username, &pool.password).await?; + let mut con = connect_cluster_node(pool, &endpoint).await?; flush_db(&mut con).await?; } Ok(()) @@ -720,9 +862,7 @@ pub async fn flush_cluster(pool: &RedisClusterPool) -> Result<(), String> { async fn cluster_total_keys(pool: &RedisClusterPool, master_nodes: &[RedisNodeEndpoint]) -> u64 { let mut total = 0; for endpoint in master_nodes { - let Ok(mut con) = - connect_direct_node(endpoint, pool.tls, pool.tls_insecure, &pool.username, &pool.password).await - else { + let Ok(mut con) = connect_cluster_node(pool, endpoint).await else { continue; }; total += redis::cmd("DBSIZE").query_async::(&mut con).await.unwrap_or(0); @@ -730,16 +870,27 @@ async fn cluster_total_keys(pool: &RedisClusterPool, master_nodes: &[RedisNodeEn total } -async fn cluster_master_nodes_from_seeds( +async fn cluster_slot_ranges_from_seeds( seed_nodes: &[RedisNodeEndpoint], tls: bool, insecure: bool, username: &str, password: &str, -) -> Result, String> { +) -> Result, String> { + let seed_routes = identity_routes(seed_nodes); + cluster_slot_ranges_from_routes(&seed_routes, tls, insecure, username, password).await +} + +async fn cluster_slot_ranges_from_routes( + seed_routes: &[RedisNodeRoute], + tls: bool, + insecure: bool, + username: &str, + password: &str, +) -> Result, String> { let mut last_error = None; - for endpoint in seed_nodes { - let mut con = match connect_direct_node(endpoint, tls, insecure, username, password).await { + for route in seed_routes { + let mut con = match connect_direct_node(&route.connect, tls, insecure, username, password).await { Ok(con) => con, Err(err) => { last_error = Some(err); @@ -753,22 +904,21 @@ async fn cluster_master_nodes_from_seeds( continue; } }; - let nodes = parse_cluster_slots(raw, &endpoint.host)?; - if !nodes.is_empty() { - return Ok(nodes); + let slot_ranges = parse_cluster_slots(raw, &route.advertised.host)?; + if !slot_ranges.is_empty() { + return Ok(slot_ranges); } } Err(last_error.unwrap_or_else(|| "Redis cluster master discovery failed".to_string())) } -fn parse_cluster_slots(raw: RedisRawValue, fallback_host: &str) -> Result, String> { +fn parse_cluster_slots(raw: RedisRawValue, fallback_host: &str) -> Result, String> { let RedisRawValue::Array(slots) = raw else { return Err("Invalid Redis CLUSTER SLOTS response".to_string()); }; - let mut seen = std::collections::HashSet::new(); - let mut nodes = Vec::new(); + let mut slot_ranges = Vec::new(); for slot in slots { let RedisRawValue::Array(parts) = slot else { continue; @@ -776,14 +926,20 @@ fn parse_cluster_slots(raw: RedisRawValue, fallback_host: &str) -> Result Result, String> { @@ -807,6 +963,111 @@ fn parse_cluster_slot_master(value: RedisRawValue, fallback_host: &str) -> Resul Ok(Some(RedisNodeEndpoint { host, port })) } +fn parse_cluster_slot_number(value: &str) -> Result { + let slot = value.parse::().map_err(|_| format!("Invalid Redis cluster slot '{value}'"))?; + if slot > 16_383 { + return Err(format!("Invalid Redis cluster slot '{value}'")); + } + Ok(slot) +} + +fn identity_routes(endpoints: &[RedisNodeEndpoint]) -> Vec { + endpoints + .iter() + .cloned() + .map(|endpoint| RedisNodeRoute { advertised: endpoint.clone(), connect: endpoint }) + .collect() +} + +pub fn unique_master_nodes(slot_ranges: &[RedisClusterSlotRange]) -> Vec { + let mut seen = std::collections::HashSet::new(); + let mut nodes = Vec::new(); + for slot_range in slot_ranges { + if seen.insert((slot_range.master.host.clone(), slot_range.master.port)) { + nodes.push(slot_range.master.clone()); + } + } + nodes +} + +pub async fn cluster_any_connection(pool: &RedisClusterPool) -> Result, String> { + if let Some(connection) = &pool.connection { + return Ok(RedisClusterConnectionGuard::Native(connection.lock().await)); + } + let endpoint = pool + .node_routes + .first() + .or_else(|| pool.seed_routes.first()) + .map(|route| &route.advertised) + .ok_or_else(|| "Redis cluster has no routable nodes".to_string())?; + connect_cluster_node(pool, endpoint).await.map(RedisClusterConnectionGuard::Direct) +} + +pub async fn cluster_key_connection<'a>( + pool: &'a RedisClusterPool, + key: &[u8], +) -> Result, String> { + if let Some(connection) = &pool.connection { + return Ok(RedisClusterConnectionGuard::Native(connection.lock().await)); + } + let endpoint = cluster_master_for_key(pool, key)?; + connect_cluster_node(pool, &endpoint).await.map(RedisClusterConnectionGuard::Direct) +} + +async fn connect_cluster_node( + pool: &RedisClusterPool, + advertised_endpoint: &RedisNodeEndpoint, +) -> Result { + let connect_endpoint = mapped_cluster_endpoint(pool, advertised_endpoint); + connect_direct_node(&connect_endpoint, pool.tls, pool.tls_insecure, &pool.username, &pool.password).await +} + +fn mapped_cluster_endpoint(pool: &RedisClusterPool, advertised_endpoint: &RedisNodeEndpoint) -> RedisNodeEndpoint { + pool.node_routes + .iter() + .chain(pool.seed_routes.iter()) + .find(|route| route.advertised == *advertised_endpoint) + .map(|route| route.connect.clone()) + .unwrap_or_else(|| advertised_endpoint.clone()) +} + +fn cluster_master_for_key(pool: &RedisClusterPool, key: &[u8]) -> Result { + let slot = redis_cluster_slot(key); + pool.slot_ranges + .iter() + .find(|range| range.start <= slot && slot <= range.end) + .map(|range| range.master.clone()) + .ok_or_else(|| format!("Redis cluster slot {slot} has no known master")) +} + +fn redis_cluster_slot(key: &[u8]) -> u16 { + let hashtag = key.iter().position(|byte| *byte == b'{').and_then(|start| { + key[start + 1..].iter().position(|byte| *byte == b'}').and_then(|relative_end| { + if relative_end == 0 { + None + } else { + Some(&key[start + 1..start + 1 + relative_end]) + } + }) + }); + crc16_xmodem(hashtag.unwrap_or(key)) % 16_384 +} + +fn crc16_xmodem(bytes: &[u8]) -> u16 { + let mut crc = 0_u16; + for byte in bytes { + crc ^= (*byte as u16) << 8; + for _ in 0..8 { + if (crc & 0x8000) != 0 { + crc = (crc << 1) ^ 0x1021; + } else { + crc <<= 1; + } + } + } + crc +} + pub fn parse_command_argv(command_text: &str) -> Result, String> { // Strip trailing semicolons so commands like "HGETALL aaa;" work naturally let command_text = command_text.trim_end().trim_end_matches(';'); @@ -1833,11 +2094,11 @@ mod tests { use super::{ classify_command, connection_info, decode_cluster_cursor, encode_cluster_cursor, is_redis_json_type, parse_cluster_slots, parse_command_argv, parse_database_count, parse_redis_endpoint, parse_scan_keys, - parse_stream_entries, redis_auth_candidates, redis_command_raw_to_json, redis_database_index, - redis_json_raw_to_json, redis_json_value_preview, redis_key_bytes_to_display, redis_key_bytes_to_raw, - redis_key_matches_query, redis_key_raw_to_bytes, redis_key_value_preview, redis_raw_to_json, - redis_value_contains_binary, redis_value_matches_query, RedisAuthCandidate, RedisCommandSafety, - RedisNodeEndpoint, RedisRawValue, + parse_stream_entries, redis_auth_candidates, redis_cluster_slot, redis_command_raw_to_json, + redis_database_index, redis_json_raw_to_json, redis_json_value_preview, redis_key_bytes_to_display, + redis_key_bytes_to_raw, redis_key_matches_query, redis_key_raw_to_bytes, redis_key_value_preview, + redis_raw_to_json, redis_value_contains_binary, redis_value_matches_query, RedisAuthCandidate, + RedisClusterSlotRange, RedisCommandSafety, RedisNodeEndpoint, RedisRawValue, }; use crate::models::connection::ConnectionConfig; use redis::ConnectionAddr; @@ -2200,12 +2461,26 @@ mod tests { assert_eq!( parse_cluster_slots(raw, "127.0.0.1").unwrap(), vec![ - RedisNodeEndpoint { host: "10.0.0.1".to_string(), port: 7000 }, - RedisNodeEndpoint { host: "10.0.0.2".to_string(), port: 7001 }, + RedisClusterSlotRange { + start: 0, + end: 5460, + master: RedisNodeEndpoint { host: "10.0.0.1".to_string(), port: 7000 }, + }, + RedisClusterSlotRange { + start: 5461, + end: 10922, + master: RedisNodeEndpoint { host: "10.0.0.2".to_string(), port: 7001 }, + }, ] ); } + #[test] + fn calculates_redis_cluster_hash_tag_slots() { + assert_eq!(redis_cluster_slot(b"issue1246:{user}:a"), redis_cluster_slot(b"issue1246:{user}:b")); + assert_ne!(redis_cluster_slot(b"issue1246:{user}:a"), redis_cluster_slot(b"issue1246:{other}:a")); + } + #[test] fn parses_redis_json_get_bulk_string() { let raw = bulk(r#"{"id":1,"embedding":[0.1,0.2],"meta":{"source":"test"}}"#); diff --git a/crates/dbx-core/src/db/ssh_tunnel.rs b/crates/dbx-core/src/db/ssh_tunnel.rs index 989bb4426..608921eab 100644 --- a/crates/dbx-core/src/db/ssh_tunnel.rs +++ b/crates/dbx-core/src/db/ssh_tunnel.rs @@ -679,6 +679,18 @@ impl TunnelManager { } } } + + pub async fn stop_tunnels_with_prefix(&self, connection_id_prefix: &str) { + let mut tunnels = self.tunnels.lock().await; + let keys: Vec = tunnels.keys().filter(|key| key.starts_with(connection_id_prefix)).cloned().collect(); + for key in keys { + if let Some(entry) = tunnels.remove(&key) { + for handle in entry.handles { + handle.abort(); + } + } + } + } } #[allow(clippy::too_many_arguments)] diff --git a/crates/dbx-core/src/redis_ops.rs b/crates/dbx-core/src/redis_ops.rs index 311e84754..e4a13bf5b 100644 --- a/crates/dbx-core/src/redis_ops.rs +++ b/crates/dbx-core/src/redis_ops.rs @@ -144,8 +144,8 @@ pub async fn redis_get_value_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::get_value(&mut *con, &key).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::get_value(&mut con, &key).await } } } @@ -185,8 +185,8 @@ pub async fn redis_set_string_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::set_string(&mut *con, &key, value, ttl).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::set_string(&mut con, &key, value, ttl).await } } } @@ -218,8 +218,8 @@ pub async fn redis_delete_key_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::delete_key(&mut *con, &key).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::delete_key(&mut con, &key).await } } } @@ -260,8 +260,8 @@ pub async fn redis_hash_set_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::hash_set(&mut *con, &key, field, value, ttl).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::hash_set(&mut con, &key, field, value, ttl).await } } } @@ -293,8 +293,8 @@ pub async fn redis_hash_del_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::hash_del(&mut *con, &key, field).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::hash_del(&mut con, &key, field).await } } } @@ -333,8 +333,8 @@ pub async fn redis_list_push_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::list_push(&mut *con, &key, value, ttl).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::list_push(&mut con, &key, value, ttl).await } } } @@ -363,8 +363,8 @@ pub async fn redis_list_set_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::list_set(&mut *con, &key, index, value).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::list_set(&mut con, &key, index, value).await } } } @@ -401,8 +401,8 @@ pub async fn redis_list_remove_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::list_remove(&mut *con, &key, index).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::list_remove(&mut con, &key, index).await } } } @@ -441,8 +441,8 @@ pub async fn redis_set_add_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::set_add(&mut *con, &key, member, ttl).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::set_add(&mut con, &key, member, ttl).await } } } @@ -479,8 +479,8 @@ pub async fn redis_set_remove_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::set_remove(&mut *con, &key, member).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::set_remove(&mut con, &key, member).await } } } @@ -510,8 +510,8 @@ pub async fn redis_zadd_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::zadd(&mut *con, &key, member, score, ttl).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::zadd(&mut con, &key, member, score, ttl).await } } } @@ -539,8 +539,8 @@ pub async fn redis_zrem_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::zrem(&mut *con, &key, member).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::zrem(&mut con, &key, member).await } } } @@ -570,8 +570,8 @@ pub async fn redis_stream_add_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::stream_add(&mut *con, &key, entry_id, &fields, ttl).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::stream_add(&mut con, &key, entry_id, &fields, ttl).await } } } @@ -600,8 +600,8 @@ pub async fn redis_json_set_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::json_set(&mut *con, &key, value, ttl).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::json_set(&mut con, &key, value, ttl).await } } } @@ -625,8 +625,8 @@ pub async fn redis_check_json_module_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::check_json_module(&mut *con).await + let mut con = redis_driver::cluster_any_connection(cluster).await?; + redis_driver::check_json_module(&mut con).await } }, _ => Err("Not a Redis connection".to_string()), @@ -653,8 +653,8 @@ pub async fn redis_set_ttl_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::set_ttl(&mut *con, &key, ttl).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::set_ttl(&mut con, &key, ttl).await } } } @@ -683,10 +683,10 @@ pub async fn redis_delete_keys_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; let mut deleted = 0; for key in &keys { - deleted += redis_driver::delete_keys(&mut *con, std::slice::from_ref(key)).await?; + let mut con = redis_driver::cluster_key_connection(cluster, key).await?; + deleted += redis_driver::delete_keys(&mut con, std::slice::from_ref(key)).await?; } Ok(deleted) } @@ -733,19 +733,39 @@ pub async fn redis_execute_command_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - if let Ok(argv) = redis_driver::parse_command_argv(command) { - if argv.first().is_some_and(|name| name.eq_ignore_ascii_case("SELECT")) { - return Err("Redis Cluster only supports db0; SELECT is not available".to_string()); + let mut con = if let Ok(argv) = redis_driver::parse_command_argv(command) { + if let Some(command_name) = argv.first() { + if command_name.eq_ignore_ascii_case("SELECT") { + return Err("Redis Cluster only supports db0; SELECT is not available".to_string()); + } } - } - let mut con = cluster.connection.lock().await; - redis_driver::execute_command(&mut *con, command, skip_safety_check).await + if command_may_target_first_key(&argv) { + redis_driver::cluster_key_connection(cluster, argv[1].as_bytes()).await? + } else { + redis_driver::cluster_any_connection(cluster).await? + } + } else { + redis_driver::cluster_any_connection(cluster).await? + }; + redis_driver::execute_command(&mut con, command, skip_safety_check).await } }, _ => Err("Not a Redis connection".to_string()), } } +fn command_may_target_first_key(argv: &[String]) -> bool { + if argv.len() < 2 { + return false; + } + match argv[0].to_ascii_uppercase().as_str() { + "PING" | "INFO" | "DBSIZE" | "TIME" | "ROLE" | "CLUSTER" | "CLIENT" | "COMMAND" | "HELLO" | "AUTH" | "QUIT" => { + false + } + _ => true, + } +} + pub async fn redis_load_more_in_db_core( state: &AppState, connection_id: &str, @@ -768,8 +788,8 @@ pub async fn redis_load_more_in_db_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::load_more_collection(&mut *con, &key, key_type, cursor, count).await + let mut con = redis_driver::cluster_key_connection(cluster, &key).await?; + redis_driver::load_more_collection(&mut con, &key, key_type, cursor, count).await } } } @@ -795,8 +815,8 @@ pub async fn redis_publish_core( } RedisConnection::Cluster(cluster) => { redis_driver::ensure_cluster_db(db)?; - let mut con = cluster.connection.lock().await; - redis_driver::publish_message(&mut *con, channel, message).await + let mut con = redis_driver::cluster_any_connection(cluster).await?; + redis_driver::publish_message(&mut con, channel, message).await } }, _ => Err("Not a Redis connection".to_string()), diff --git a/src-tauri/src/commands/connection.rs b/src-tauri/src/commands/connection.rs index 387e3d3b8..8cf54bbd5 100644 --- a/src-tauri/src/commands/connection.rs +++ b/src-tauri/src/commands/connection.rs @@ -548,7 +548,7 @@ pub async fn test_connection(state: State<'_, Arc>, config: Connection } DatabaseType::Redis => { let con = if config.uses_redis_cluster() { - db::redis_driver::connect_cluster(&config).await?; + state.connect_redis_cluster(&tunnel_id, &config).await?; return Ok("Connection successful".to_string()); } else if config.uses_redis_sentinel() { db::redis_driver::connect_sentinel(&config).await? @@ -779,7 +779,7 @@ pub async fn connect_db(state: State<'_, Arc>, config: ConnectionConfi DatabaseType::Redis => { let con = if db_config.uses_redis_cluster() { PoolKind::Redis(db::redis_driver::RedisConnection::Cluster( - db::redis_driver::connect_cluster(&db_config).await?, + state.connect_redis_cluster(&id, &db_config).await?, )) } else if db_config.uses_redis_sentinel() { PoolKind::Redis(db::redis_driver::RedisConnection::Direct(tokio::sync::Mutex::new(