use super::get_valkey_connection_info;
use super::reconnecting_connection::{ReconnectReason, ReconnectingConnection};
use super::{ConnectionRequest, NodeAddress, NodeDiscoveryMode, TlsMode};
use crate::client::types::ReadFrom as ClientReadFrom;
use futures::{StreamExt, future, stream};
use logger_core::log_debug;
use logger_core::log_info;
use logger_core::log_warn;
use redis::aio::ConnectionLike;
use redis::cluster_routing::{self, ResponsePolicy, Routable, RoutingInfo, is_readonly_cmd};
use redis::{AddressResolver, PushInfo, RedisError, RedisResult, RetryStrategy, Value};
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::time::Duration;
use telemetrylib::Telemetry;
use tokio::sync::mpsc;
use tokio::task;
#[derive(Debug)]
enum ReadFrom {
Primary,
PreferReplica {
latest_read_replica_index: Arc<AtomicUsize>,
},
AllNodes {
latest_read_node_index: Arc<AtomicUsize>,
},
AZAffinity {
client_az: String,
last_read_replica_index: Arc<AtomicUsize>,
},
AZAffinityReplicasAndPrimary {
client_az: String,
last_read_replica_index: Arc<AtomicUsize>,
},
}
#[derive(Debug)]
struct DropWrapper {
primary_index: usize,
nodes: Vec<ReconnectingConnection>,
read_from: ReadFrom,
read_only: bool,
}
impl Drop for DropWrapper {
fn drop(&mut self) {
for node in self.nodes.iter() {
node.mark_as_dropped();
}
}
}
#[derive(Clone, Debug)]
pub struct StandaloneClient {
inner: Arc<DropWrapper>,
}
impl Drop for StandaloneClient {
fn drop(&mut self) {
Telemetry::decr_total_clients(1);
}
}
pub enum StandaloneClientConnectionError {
NoAddressesProvided,
FailedConnection(Vec<(Option<String>, RedisError)>),
PrimaryConflictFound(String),
}
impl std::fmt::Debug for StandaloneClientConnectionError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
StandaloneClientConnectionError::NoAddressesProvided => {
write!(f, "No addresses provided")
}
StandaloneClientConnectionError::FailedConnection(errs) => {
match errs.len() {
0 => {
writeln!(f, "Failed without explicit error")?;
}
1 => {
let (ref address, ref error) = errs[0];
match address {
Some(address) => {
writeln!(f, "Received error for address `{address}`: {error}")?
}
None => writeln!(f, "Received error: {error}")?,
}
}
_ => {
writeln!(f, "Received errors:")?;
for (address, error) in errs {
match address {
Some(address) => writeln!(f, "{address}: {error}")?,
None => writeln!(f, "{error}")?,
}
}
}
};
Ok(())
}
StandaloneClientConnectionError::PrimaryConflictFound(found_primaries) => {
writeln!(
f,
"Primary conflict. More than one primary found in a Standalone setup: {found_primaries}"
)
}
}
}
}
impl StandaloneClient {
pub async fn create_client(
connection_request: ConnectionRequest,
push_sender: Option<mpsc::UnboundedSender<PushInfo>>,
iam_token_manager: Option<&Arc<crate::iam::IAMTokenManager>>,
pubsub_synchronizer: Option<Arc<dyn crate::pubsub::PubSubSynchronizer>>,
) -> Result<Self, StandaloneClientConnectionError> {
if connection_request.addresses.is_empty() {
return Err(StandaloneClientConnectionError::NoAddressesProvided);
}
if connection_request.read_only
&& matches!(
connection_request.read_from,
Some(ClientReadFrom::AZAffinity(_))
| Some(ClientReadFrom::AZAffinityReplicasAndPrimary(_))
)
{
return Err(StandaloneClientConnectionError::FailedConnection(vec![(
None,
RedisError::from((
redis::ErrorKind::InvalidClientConfig,
"read-only mode is not compatible with AZAffinity strategies",
)),
)]));
}
if connection_request.read_only
&& connection_request.node_discovery_mode == NodeDiscoveryMode::DiscoverAll
{
return Err(StandaloneClientConnectionError::FailedConnection(vec![(
None,
RedisError::from((
redis::ErrorKind::InvalidClientConfig,
"read-only mode is not compatible with DISCOVER_ALL node discovery mode",
)),
)]));
}
let valkey_connection_info =
get_valkey_connection_info(&connection_request, iam_token_manager).await;
let retry_strategy = match connection_request.connection_retry_strategy {
Some(strategy) => RetryStrategy::new(
strategy.exponent_base,
strategy.factor,
strategy.number_of_retries,
strategy.jitter_percent,
),
None => RetryStrategy::default(),
};
let tls_mode = connection_request.tls_mode;
let node_count = connection_request.addresses.len();
let discover_az = matches!(
connection_request.read_from,
Some(ClientReadFrom::AZAffinity(_))
| Some(ClientReadFrom::AZAffinityReplicasAndPrimary(_))
);
let connection_timeout = connection_request.get_connection_timeout();
let tcp_nodelay = connection_request.tcp_nodelay;
let has_root_certs = !connection_request.root_certs.is_empty();
let has_client_cert = !connection_request.client_cert.is_empty();
let has_client_key = !connection_request.client_key.is_empty();
if has_client_cert != has_client_key {
return Err(StandaloneClientConnectionError::FailedConnection(vec![(
None,
RedisError::from((
redis::ErrorKind::InvalidClientConfig,
"client_cert and client_key must both be provided or both be empty",
)),
)]));
}
let tls_params = if has_root_certs || has_client_cert || has_client_key {
if tls_mode.unwrap_or(TlsMode::NoTls) == TlsMode::NoTls {
return Err(StandaloneClientConnectionError::FailedConnection(vec![(
None,
RedisError::from((
redis::ErrorKind::InvalidClientConfig,
"TLS certificates provided but TLS is disabled",
)),
)]));
}
let root_cert = if has_root_certs {
let mut combined_certs = Vec::new();
for cert in &connection_request.root_certs {
combined_certs.extend_from_slice(cert);
}
Some(combined_certs)
} else {
None
};
let client_tls = if has_client_cert && has_client_key {
Some(redis::ClientTlsConfig {
client_cert: connection_request.client_cert.clone(),
client_key: connection_request.client_key.clone(),
})
} else {
None
};
let tls_certificates = redis::TlsCertificates {
client_tls,
root_cert,
};
Some(
redis::retrieve_tls_certificates(tls_certificates).map_err(|err| {
StandaloneClientConnectionError::FailedConnection(vec![(None, err)])
})?,
)
} else {
None
};
let read_only = connection_request.read_only;
let node_discovery_mode = connection_request.node_discovery_mode;
let addresses = connection_request.addresses.clone();
let read_from_option = connection_request.read_from.clone();
let iam_token_handle = iam_token_manager.map(|m| m.get_token_handle());
let discovery_conn_info = valkey_connection_info.clone();
let discovery_push_sender = push_sender.clone();
let discovery_tls_params = tls_params.clone();
let discovery_pubsub_sync = pubsub_synchronizer.clone();
let discovery_iam_handle = iam_token_handle.clone();
let discovery_resolver = connection_request.address_resolver.clone();
let mut stream = stream::iter(addresses)
.map(move |address| {
let info = valkey_connection_info.clone();
let retry = retry_strategy;
let sender = push_sender.clone();
let tls = tls_mode.unwrap_or(TlsMode::NoTls);
let discover = discover_az;
let timeout = connection_timeout;
let params = tls_params.clone();
let nodelay = tcp_nodelay;
let sync = pubsub_synchronizer.clone();
let skip_replication =
read_only || node_discovery_mode == NodeDiscoveryMode::Static;
let resolver = connection_request.address_resolver.clone();
let iam_handle = iam_token_handle.clone();
async move {
get_connection_and_replication_info(
&address,
&retry,
&info,
tls,
&sender,
discover,
timeout,
params,
nodelay,
&sync,
skip_replication,
resolver.as_ref(),
iam_handle,
)
.await
.map_err(|err| (format!("{}:{}", address.host, address.port), err))
}
})
.buffer_unordered(node_count);
let mut nodes = Vec::with_capacity(node_count);
let mut addresses_and_errors = Vec::with_capacity(node_count);
let mut primary_index = if read_only || node_discovery_mode == NodeDiscoveryMode::Static {
Some(0)
} else {
None
};
let mut replication_infos: Vec<Option<String>> = Vec::with_capacity(node_count);
while let Some(result) = stream.next().await {
match result {
Ok((connection, replication_status)) => {
nodes.push(connection);
let info_str = replication_status
.and_then(|status| redis::from_owned_redis_value::<String>(status).ok());
let is_primary = info_str
.as_ref()
.is_some_and(|val| val.contains("role:master"));
replication_infos.push(info_str);
if is_primary {
if let Some(existing_primary) = primary_index {
return Err(StandaloneClientConnectionError::PrimaryConflictFound(
format!(
"Primary nodes: {:?}, {:?}",
nodes.pop(),
nodes.get(existing_primary)
),
));
}
primary_index = Some(nodes.len().saturating_sub(1));
}
}
Err((address, (connection, err))) => {
nodes.push(connection);
replication_infos.push(None);
addresses_and_errors.push((Some(address), err));
}
}
}
if node_discovery_mode == NodeDiscoveryMode::DiscoverAll {
let mut discovered: Vec<NodeAddress> = Vec::new();
let existing: Vec<String> = connection_request
.addresses
.iter()
.map(|a| format!("{}:{}", a.host, a.port))
.collect();
for info_str in replication_infos.iter().flatten() {
if is_primary_role(info_str) {
let replicas = parse_replica_addresses(info_str);
log_info(
"topology discovery",
format!("Discovered {} replica(s) from primary", replicas.len()),
);
for r in replicas {
if !address_is_known(&r, &existing, &discovered) {
discovered.push(r);
}
}
} else if let Some(primary_addr) = parse_primary_address(info_str) {
log_info(
"topology discovery",
format!(
"Discovered primary at {}:{}",
primary_addr.host, primary_addr.port
),
);
if !address_is_known(&primary_addr, &existing, &discovered) {
discovered.push(primary_addr);
}
}
}
let tls = tls_mode.unwrap_or(TlsMode::NoTls);
let discovered_count = discovered.len();
let mut phase2_stream = stream::iter(discovered.clone())
.map(|address| {
let conn_info = discovery_conn_info.clone();
let sender = discovery_push_sender.clone();
let params = discovery_tls_params.clone();
let sync = discovery_pubsub_sync.clone();
let iam_handle = discovery_iam_handle.clone();
let resolver = discovery_resolver.clone();
async move {
let result = get_connection_and_replication_info(
&address,
&retry_strategy,
&conn_info,
tls,
&sender,
discover_az,
connection_timeout,
params,
tcp_nodelay,
&sync,
false,
resolver.as_ref(),
iam_handle,
)
.await;
(address, result)
}
})
.buffer_unordered(discovered_count);
let mut phase3_addresses: Vec<NodeAddress> = Vec::new();
while let Some((addr, result)) = phase2_stream.next().await {
match result {
Ok((connection, replication_status)) => {
let info_str = replication_status
.and_then(|s| redis::from_owned_redis_value::<String>(s).ok());
let is_primary =
info_str.as_ref().is_some_and(|v| v.contains("role:master"));
if is_primary && primary_index.is_none() {
primary_index = Some(nodes.len());
}
nodes.push(connection);
if let Some(info) = info_str.as_deref().filter(|_| is_primary) {
for r in parse_replica_addresses(info) {
if !address_is_known(&r, &existing, &discovered)
&& !address_is_known(&r, &existing, &phase3_addresses)
{
phase3_addresses.push(r);
}
}
}
}
Err((_connection, err)) => {
log_warn(
"topology discovery",
format!(
"Failed to connect to discovered node {}:{}: {}",
addr.host, addr.port, err
),
);
}
}
}
if !phase3_addresses.is_empty() {
let phase3_count = phase3_addresses.len();
let mut phase3_stream = stream::iter(phase3_addresses)
.map(|address| {
let conn_info = discovery_conn_info.clone();
let sender = discovery_push_sender.clone();
let params = discovery_tls_params.clone();
let sync = discovery_pubsub_sync.clone();
let iam_handle = discovery_iam_handle.clone();
let resolver = discovery_resolver.clone();
async move {
let result = get_connection_and_replication_info(
&address,
&retry_strategy,
&conn_info,
tls,
&sender,
discover_az,
connection_timeout,
params,
tcp_nodelay,
&sync,
false,
resolver.as_ref(),
iam_handle,
)
.await;
(address, result)
}
})
.buffer_unordered(phase3_count);
while let Some((addr, result)) = phase3_stream.next().await {
match result {
Ok((conn, _)) => nodes.push(conn),
Err((_conn, err)) => {
log_warn(
"topology discovery",
format!(
"Failed to connect to discovered replica {}:{}: {}",
addr.host, addr.port, err
),
);
}
}
}
}
if !discovered.is_empty() {
log_info(
"topology discovery",
format!("Full topology: {} node(s) connected", nodes.len()),
);
}
}
let primary_index = if read_only {
if nodes.is_empty() && !addresses_and_errors.is_empty() {
return Err(StandaloneClientConnectionError::FailedConnection(
addresses_and_errors,
));
}
0 } else {
match primary_index {
Some(idx) => idx,
None => {
let mut errors = addresses_and_errors;
if errors.is_empty() {
errors.insert(
0,
(
None,
RedisError::from((
redis::ErrorKind::ClientError,
"No primary node found",
)),
),
)
};
return Err(StandaloneClientConnectionError::FailedConnection(errors));
}
}
};
if !addresses_and_errors.is_empty() {
log_warn(
"client creation",
format!(
"Failed to connect to {addresses_and_errors:?}, will attempt to reconnect."
),
);
}
let read_from = if read_only && read_from_option.is_none() {
ReadFrom::PreferReplica {
latest_read_replica_index: Default::default(),
}
} else {
get_read_from(read_from_option)
};
#[cfg(feature = "standalone_heartbeat")]
for node in nodes.iter() {
Self::start_heartbeat(node.clone());
}
for node in nodes.iter() {
Self::start_periodic_connection_check(node.clone());
}
Telemetry::incr_total_clients(1);
Ok(Self {
inner: Arc::new(DropWrapper {
primary_index,
nodes,
read_from,
read_only,
}),
})
}
fn get_primary_connection(&self) -> &ReconnectingConnection {
self.inner.nodes.get(self.inner.primary_index).unwrap()
}
fn round_robin_read_from_replica(
&self,
latest_read_replica_index: &Arc<AtomicUsize>,
) -> &ReconnectingConnection {
let initial_index = latest_read_replica_index.load(Ordering::Relaxed);
let mut check_count = 0;
loop {
check_count += 1;
if check_count > self.inner.nodes.len() {
return self.get_primary_connection();
}
let index = (initial_index + check_count) % self.inner.nodes.len();
if index == self.inner.primary_index {
continue;
}
let Some(connection) = self.inner.nodes.get(index) else {
continue;
};
if connection.is_connected() {
let _ = latest_read_replica_index.compare_exchange_weak(
initial_index,
index,
Ordering::Relaxed,
Ordering::Relaxed,
);
return connection;
}
}
}
fn round_robin_read_from_all_nodes(
&self,
latest_read_node_index: &Arc<AtomicUsize>,
) -> &ReconnectingConnection {
let initial_index = latest_read_node_index.load(Ordering::Relaxed);
let mut check_count = 0;
loop {
check_count += 1;
if check_count > self.inner.nodes.len() {
return self.get_primary_connection();
}
let index = (initial_index + check_count) % self.inner.nodes.len();
let Some(connection) = self.inner.nodes.get(index) else {
continue;
};
if connection.is_connected() {
let _ = latest_read_node_index.compare_exchange_weak(
initial_index,
index,
Ordering::Relaxed,
Ordering::Relaxed,
);
return connection;
}
}
}
async fn get_next_local_replica(
&self,
latest_read_replica_index: &Arc<AtomicUsize>,
client_az: &str,
) -> Option<&ReconnectingConnection> {
let initial_index = latest_read_replica_index.load(Ordering::Relaxed);
let mut retries = 0usize;
loop {
retries = retries.saturating_add(1);
if retries > self.inner.nodes.len() {
return None;
}
let index = (initial_index + retries) % self.inner.nodes.len();
if index == self.inner.primary_index {
continue;
}
let replica = &self.inner.nodes[index];
if let Ok(connection) = replica.get_connection().await
&& let Some(replica_az) = connection.get_az().as_deref()
&& replica_az == client_az
{
let _ = latest_read_replica_index.compare_exchange_weak(
initial_index,
index,
Ordering::Relaxed,
Ordering::Relaxed,
);
return Some(replica);
}
}
}
async fn round_robin_read_from_replica_az_awareness(
&self,
latest_read_replica_index: &Arc<AtomicUsize>,
client_az: &str,
) -> &ReconnectingConnection {
if let Some(replica) = self
.get_next_local_replica(latest_read_replica_index, client_az)
.await
{
return replica;
}
self.round_robin_read_from_replica(latest_read_replica_index)
}
async fn round_robin_read_from_replica_az_awareness_replicas_and_primary(
&self,
latest_read_replica_index: &Arc<AtomicUsize>,
client_az: &str,
) -> &ReconnectingConnection {
if let Some(replica) = self
.get_next_local_replica(latest_read_replica_index, client_az)
.await
{
return replica;
}
let primary = self.get_primary_connection();
if let Ok(connection) = primary.get_connection().await
&& let Some(primary_az) = connection.get_az().as_deref()
&& primary_az == client_az
{
return primary;
}
self.round_robin_read_from_all_nodes(latest_read_replica_index)
}
async fn get_connection(&self, readonly: bool) -> &ReconnectingConnection {
if self.inner.nodes.len() == 1 || !readonly {
return self.get_primary_connection();
}
match &self.inner.read_from {
ReadFrom::Primary => self.get_primary_connection(),
ReadFrom::PreferReplica {
latest_read_replica_index,
} => self.round_robin_read_from_replica(latest_read_replica_index),
ReadFrom::AllNodes {
latest_read_node_index,
} => self.round_robin_read_from_all_nodes(latest_read_node_index),
ReadFrom::AZAffinity {
client_az,
last_read_replica_index,
} => {
self.round_robin_read_from_replica_az_awareness(last_read_replica_index, client_az)
.await
}
ReadFrom::AZAffinityReplicasAndPrimary {
client_az,
last_read_replica_index,
} => {
self.round_robin_read_from_replica_az_awareness_replicas_and_primary(
last_read_replica_index,
client_az,
)
.await
}
}
}
async fn send_request(
cmd: &redis::Cmd,
reconnecting_connection: &ReconnectingConnection,
) -> RedisResult<Value> {
cmd.watchdog_phase
.store(redis::PHASE_SENT, std::sync::atomic::Ordering::Release);
let mut connection = reconnecting_connection.get_connection().await?;
let result = connection.send_packed_command(cmd).await;
match result {
Err(err) if err.is_unrecoverable_error() => {
log_warn("send request", format!("received disconnect error `{err}`"));
reconnecting_connection.reconnect(ReconnectReason::ConnectionDropped);
Err(err)
}
_ => result,
}
}
pub(crate) async fn send_request_to_all_nodes(
&mut self,
cmd: &redis::Cmd,
response_policy: Option<ResponsePolicy>,
) -> RedisResult<Value> {
let requests = self
.inner
.nodes
.iter()
.map(|node| Self::send_request(cmd, node));
match response_policy {
Some(ResponsePolicy::AllSucceeded) => {
future::try_join_all(requests)
.await
.map(|mut results| results.pop().unwrap()) }
Some(ResponsePolicy::OneSucceeded) => future::select_ok(requests.map(Box::pin))
.await
.map(|(result, _)| result),
Some(ResponsePolicy::FirstSucceededNonEmptyOrAllEmpty) => {
future::select_ok(requests.map(|request| {
Box::pin(async move {
let result = request.await?;
match result {
Value::Nil => {
Err((redis::ErrorKind::ResponseError, "no value found").into())
}
_ => Ok(result),
}
})
}))
.await
.map(|(result, _)| result)
}
Some(ResponsePolicy::Aggregate(op)) => future::try_join_all(requests)
.await
.and_then(|results| cluster_routing::aggregate(results, op)),
Some(ResponsePolicy::AggregateArray(op)) => future::try_join_all(requests)
.await
.and_then(|results| cluster_routing::aggregate_array(results, op)),
Some(ResponsePolicy::AggregateLogical(op)) => future::try_join_all(requests)
.await
.and_then(|results| cluster_routing::logical_aggregate(results, op)),
Some(ResponsePolicy::CombineArrays) => future::try_join_all(requests)
.await
.and_then(cluster_routing::combine_array_results),
Some(ResponsePolicy::CombineMaps) => future::try_join_all(requests)
.await
.and_then(cluster_routing::combine_map_results),
Some(ResponsePolicy::Special) => {
let results = future::try_join_all(requests).await?;
let node_result_pairs = self
.inner
.nodes
.iter()
.zip(results)
.map(|(node, result)| (Value::BulkString(node.node_address().into()), result))
.collect();
Ok(Value::Map(node_result_pairs))
}
None => {
future::try_join_all(requests).await.map(Value::Array)
}
}
}
async fn send_request_to_single_node(
&mut self,
cmd: &redis::Cmd,
readonly: bool,
) -> RedisResult<Value> {
let reconnecting_connection = self.get_connection(readonly).await;
Self::send_request(cmd, reconnecting_connection).await
}
pub async fn send_command(&mut self, cmd: &redis::Cmd) -> RedisResult<Value> {
let Some(cmd_bytes) = Routable::command(cmd) else {
return self.send_request_to_single_node(cmd, false).await;
};
if self.inner.read_only && !is_readonly_cmd(cmd_bytes.as_slice()) {
return Err(RedisError::from((
redis::ErrorKind::ReadOnly,
"write commands are not allowed in read-only mode",
)));
}
if RoutingInfo::is_all_nodes(cmd_bytes.as_slice()) {
let response_policy = ResponsePolicy::for_command(cmd_bytes.as_slice());
return self.send_request_to_all_nodes(cmd, response_policy).await;
}
self.send_request_to_single_node(cmd, is_readonly_cmd(cmd_bytes.as_slice()))
.await
}
pub async fn send_pipeline(
&mut self,
pipeline: &redis::Pipeline,
offset: usize,
count: usize,
) -> RedisResult<Vec<Value>> {
let reconnecting_connection = self.get_primary_connection();
let mut connection = reconnecting_connection.get_connection().await?;
let result = connection
.send_packed_commands(pipeline, offset, count)
.await;
match result {
Err(err) if err.is_unrecoverable_error() => {
log_warn(
"pipeline request",
format!("received disconnect error `{err}`"),
);
reconnecting_connection.reconnect(ReconnectReason::ConnectionDropped);
Err(err)
}
_ => result,
}
}
#[cfg(feature = "standalone_heartbeat")]
fn start_heartbeat(reconnecting_connection: ReconnectingConnection) {
task::spawn(async move {
loop {
tokio::time::sleep(super::HEARTBEAT_SLEEP_DURATION).await;
if reconnecting_connection.is_dropped() {
log_debug(
"StandaloneClient",
"heartbeat stopped after connection was dropped",
);
return;
}
let Some(mut connection) = reconnecting_connection.try_get_connection().await
else {
log_debug(
"StandaloneClient",
"heartbeat stopped while connection is reconnecting",
);
continue;
};
log_debug("StandaloneClient", "performing heartbeat");
if connection
.send_packed_command(&redis::cmd("PING"))
.await
.is_err_and(|err| err.is_connection_dropped() || err.is_connection_refusal())
{
log_debug("StandaloneClient", "heartbeat triggered reconnect");
reconnecting_connection.reconnect(ReconnectReason::ConnectionDropped);
}
}
});
}
fn start_periodic_connection_check(reconnecting_connection: ReconnectingConnection) {
task::spawn(async move {
loop {
reconnecting_connection
.wait_for_disconnect_with_timeout(&super::CONNECTION_CHECKS_INTERVAL)
.await;
if reconnecting_connection.is_dropped() {
log_debug(
"StandaloneClient",
"connection checker stopped after connection was dropped",
);
return;
}
let Some(connection) = reconnecting_connection.try_get_connection().await else {
log_debug(
"StandaloneClient",
"connection checker is skipping a connections since its reconnecting",
);
continue;
};
if connection.is_closed() {
log_debug(
"StandaloneClient",
"connection checker has triggered reconnect",
);
reconnecting_connection.reconnect(ReconnectReason::ConnectionDropped);
}
}
});
}
pub async fn update_connection_password(
&self,
new_password: Option<String>,
) -> RedisResult<Value> {
for node in self.inner.nodes.iter() {
node.update_connection_password(new_password.clone());
}
Ok(Value::Okay)
}
pub async fn update_connection_database(&self, database_id: i64) -> RedisResult<Value> {
for node in self.inner.nodes.iter() {
node.update_connection_database(database_id);
}
Ok(Value::Okay)
}
pub async fn update_connection_client_name(
&self,
new_client_name: Option<String>,
) -> RedisResult<Value> {
for node in self.inner.nodes.iter() {
node.update_connection_client_name(new_client_name.clone());
}
Ok(Value::Okay)
}
pub async fn update_connection_username(
&self,
new_username: Option<String>,
) -> RedisResult<Value> {
for node in self.inner.nodes.iter() {
node.update_connection_username(new_username.clone());
}
Ok(Value::Okay)
}
pub async fn update_connection_protocol(
&self,
new_protocol: redis::ProtocolVersion,
) -> RedisResult<Value> {
for node in self.inner.nodes.iter() {
node.update_connection_protocol(new_protocol);
}
Ok(Value::Okay)
}
pub fn get_username(&self) -> Option<String> {
self.get_primary_connection().get_username()
}
}
#[allow(clippy::too_many_arguments)]
async fn get_connection_and_replication_info(
address: &NodeAddress,
retry_strategy: &RetryStrategy,
connection_info: &redis::RedisConnectionInfo,
tls_mode: TlsMode,
push_sender: &Option<mpsc::UnboundedSender<PushInfo>>,
discover_az: bool,
connection_timeout: Duration,
tls_params: Option<redis::TlsConnParams>,
tcp_nodelay: bool,
pubsub_synchronizer: &Option<Arc<dyn crate::pubsub::PubSubSynchronizer>>,
skip_replication_check: bool,
address_resolver: Option<&Arc<dyn AddressResolver>>,
iam_token_handle: Option<super::IAMTokenHandle>,
) -> Result<(ReconnectingConnection, Option<Value>), (ReconnectingConnection, RedisError)> {
let reconnecting_connection = ReconnectingConnection::new(
address,
*retry_strategy,
connection_info.clone(),
tls_mode,
push_sender.clone(),
discover_az,
connection_timeout,
tls_params,
tcp_nodelay,
pubsub_synchronizer.clone(),
address_resolver,
iam_token_handle,
)
.await?;
let mut multiplexed_connection = match reconnecting_connection.get_connection().await {
Ok(multiplexed_connection) => multiplexed_connection,
Err(err) => {
reconnecting_connection.reconnect(ReconnectReason::ConnectionDropped);
return Err((reconnecting_connection, err));
}
};
if skip_replication_check {
return Ok((reconnecting_connection, None));
}
match multiplexed_connection
.send_packed_command(redis::cmd("INFO").arg("REPLICATION"))
.await
{
Ok(replication_status) => Ok((reconnecting_connection, Some(replication_status))),
Err(err) => Err((reconnecting_connection, err)),
}
}
fn get_read_from(read_from: Option<super::ReadFrom>) -> ReadFrom {
match read_from {
Some(super::ReadFrom::Primary) => ReadFrom::Primary,
Some(super::ReadFrom::PreferReplica) => ReadFrom::PreferReplica {
latest_read_replica_index: Default::default(),
},
Some(super::ReadFrom::AllNodes) => ReadFrom::AllNodes {
latest_read_node_index: Default::default(),
},
Some(super::ReadFrom::AZAffinity(az)) => ReadFrom::AZAffinity {
client_az: az,
last_read_replica_index: Default::default(),
},
Some(super::ReadFrom::AZAffinityReplicasAndPrimary(az)) => {
ReadFrom::AZAffinityReplicasAndPrimary {
client_az: az,
last_read_replica_index: Default::default(),
}
}
None => ReadFrom::Primary,
}
}
fn parse_replica_addresses(replication_info: &str) -> Vec<NodeAddress> {
let mut replicas = Vec::new();
for line in replication_info.lines() {
let line = line.trim();
if !line.starts_with("slave") || !line.contains(":ip=") {
continue;
}
let after_colon = match line.split_once(':') {
Some((_, rest)) => rest,
None => continue,
};
let mut host = None;
let mut port = None;
for part in after_colon.split(',') {
if let Some(val) = part.strip_prefix("ip=") {
host = Some(val.to_string());
} else if let Some(val) = part.strip_prefix("port=") {
port = val.parse::<u16>().ok();
}
}
if let (Some(h), Some(p)) = (host, port) {
replicas.push(NodeAddress { host: h, port: p });
}
}
replicas
}
fn parse_primary_address(replication_info: &str) -> Option<NodeAddress> {
let mut host = None;
let mut port = None;
for line in replication_info.lines() {
let line = line.trim();
if let Some(val) = line.strip_prefix("master_host:") {
host = Some(val.to_string());
} else if let Some(val) = line.strip_prefix("master_port:") {
port = val.parse::<u16>().ok();
}
}
match (host, port) {
(Some(h), Some(p)) => Some(NodeAddress { host: h, port: p }),
_ => None,
}
}
fn is_primary_role(replication_info: &str) -> bool {
replication_info.lines().any(|l| l.trim() == "role:master")
}
fn address_is_known(addr: &NodeAddress, existing: &[String], discovered: &[NodeAddress]) -> bool {
let key = format!("{}:{}", addr.host, addr.port);
existing.contains(&key)
|| discovered
.iter()
.any(|a| format!("{}:{}", a.host, a.port) == key)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_replica_addresses_basic() {
let info = "role:master\nconnected_slaves:2\nslave0:ip=10.0.0.1,port=6379,state=online,offset=100,lag=0\nslave1:ip=10.0.0.2,port=6380,state=online,offset=100,lag=1\n";
let replicas = parse_replica_addresses(info);
assert_eq!(replicas.len(), 2);
assert_eq!(replicas[0].host, "10.0.0.1");
assert_eq!(replicas[0].port, 6379);
assert_eq!(replicas[1].host, "10.0.0.2");
assert_eq!(replicas[1].port, 6380);
}
#[test]
fn test_parse_replica_addresses_with_type_field() {
let info = "slave0:ip=10.0.0.1,port=6379,state=online,offset=100,lag=0,type=replica\n";
let replicas = parse_replica_addresses(info);
assert_eq!(replicas.len(), 1);
assert_eq!(replicas[0].host, "10.0.0.1");
assert_eq!(replicas[0].port, 6379);
}
#[test]
fn test_parse_replica_addresses_empty() {
let info = "role:master\nconnected_slaves:0\n";
let replicas = parse_replica_addresses(info);
assert!(replicas.is_empty());
}
#[test]
fn test_parse_primary_address() {
let info = "role:slave\nmaster_host:10.0.0.1\nmaster_port:6379\nmaster_link_status:up\n";
let primary = parse_primary_address(info);
assert!(primary.is_some());
let addr = primary.unwrap();
assert_eq!(addr.host, "10.0.0.1");
assert_eq!(addr.port, 6379);
}
#[test]
fn test_parse_primary_address_missing() {
let info = "role:master\nconnected_slaves:0\n";
assert!(parse_primary_address(info).is_none());
}
#[test]
fn test_is_primary_role() {
assert!(is_primary_role("role:master\nconnected_slaves:0\n"));
assert!(!is_primary_role("role:slave\nmaster_host:10.0.0.1\n"));
}
#[test]
fn test_parse_replica_addresses_real_world() {
let info = "# Replication\nrole:master\nconnected_slaves:2\nslave0:ip=YYY.YYY.YYY.YYY,port=6379,state=online,offset=1156932007140,lag=0,type=replica\nslave1:ip=ZZZ.ZZZ.ZZZ.ZZZ,port=6379,state=online,offset=1156932007140,lag=1,type=replica\nmaster_replid:070023374a903a57e473b41ff2fbcc2fcd06a01a\n";
let replicas = parse_replica_addresses(info);
assert_eq!(replicas.len(), 2);
assert_eq!(replicas[0].host, "YYY.YYY.YYY.YYY");
assert_eq!(replicas[1].host, "ZZZ.ZZZ.ZZZ.ZZZ");
}
#[test]
fn test_parse_primary_address_real_world() {
let info = "# Replication\nrole:slave\nmaster_host:XXX.XXX.XXX.XXX\nmaster_port:6379\nmaster_link_status:up\nmaster_last_io_seconds_ago:0\n";
let primary = parse_primary_address(info).unwrap();
assert_eq!(primary.host, "XXX.XXX.XXX.XXX");
assert_eq!(primary.port, 6379);
}
}