use super::{AsyncPushSender, RedisFuture};
use crate::{
aio::{check_resp3, ConnectionLike, MultiplexedConnection, Runtime},
cmd,
types::{RedisError, RedisResult, Value},
AsyncConnectionConfig, Client, Cmd, ToRedisArgs,
};
use arc_swap::ArcSwap;
use backon::{ExponentialBuilder, Retryable};
use futures::{
future::{self, Shared},
FutureExt,
};
use futures_util::future::BoxFuture;
use std::sync::Arc;
#[derive(Clone)]
pub struct ConnectionManagerConfig {
exponent_base: u64,
factor: u64,
number_of_retries: usize,
max_delay: Option<u64>,
response_timeout: Option<std::time::Duration>,
connection_timeout: Option<std::time::Duration>,
push_sender: Option<Arc<dyn AsyncPushSender>>,
}
impl std::fmt::Debug for ConnectionManagerConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> Result<(), std::fmt::Error> {
let &Self {
exponent_base,
factor,
number_of_retries,
max_delay,
response_timeout,
connection_timeout,
push_sender,
} = &self;
f.debug_struct("ConnectionManagerConfig")
.field("exponent_base", &exponent_base)
.field("factor", &factor)
.field("number_of_retries", &number_of_retries)
.field("max_delay", &max_delay)
.field("response_timeout", &response_timeout)
.field("connection_timeout", &connection_timeout)
.field(
"push_sender",
if push_sender.is_some() {
&"set"
} else {
&"not set"
},
)
.finish()
}
}
impl ConnectionManagerConfig {
const DEFAULT_CONNECTION_RETRY_EXPONENT_BASE: u64 = 2;
const DEFAULT_CONNECTION_RETRY_FACTOR: u64 = 100;
const DEFAULT_NUMBER_OF_CONNECTION_RETRIES: usize = 6;
const DEFAULT_RESPONSE_TIMEOUT: Option<std::time::Duration> = None;
const DEFAULT_CONNECTION_TIMEOUT: Option<std::time::Duration> = None;
pub fn new() -> Self {
Self::default()
}
pub fn set_factor(mut self, factor: u64) -> ConnectionManagerConfig {
self.factor = factor;
self
}
pub fn set_max_delay(mut self, time: u64) -> ConnectionManagerConfig {
self.max_delay = Some(time);
self
}
pub fn set_exponent_base(mut self, base: u64) -> ConnectionManagerConfig {
self.exponent_base = base;
self
}
pub fn set_number_of_retries(mut self, amount: usize) -> ConnectionManagerConfig {
self.number_of_retries = amount;
self
}
pub fn set_response_timeout(
mut self,
duration: std::time::Duration,
) -> ConnectionManagerConfig {
self.response_timeout = Some(duration);
self
}
pub fn set_connection_timeout(
mut self,
duration: std::time::Duration,
) -> ConnectionManagerConfig {
self.connection_timeout = Some(duration);
self
}
pub fn set_push_sender(mut self, sender: impl AsyncPushSender) -> Self {
self.push_sender = Some(Arc::new(sender));
self
}
}
impl Default for ConnectionManagerConfig {
fn default() -> Self {
Self {
exponent_base: Self::DEFAULT_CONNECTION_RETRY_EXPONENT_BASE,
factor: Self::DEFAULT_CONNECTION_RETRY_FACTOR,
number_of_retries: Self::DEFAULT_NUMBER_OF_CONNECTION_RETRIES,
max_delay: None,
response_timeout: Self::DEFAULT_RESPONSE_TIMEOUT,
connection_timeout: Self::DEFAULT_CONNECTION_TIMEOUT,
push_sender: None,
}
}
}
#[derive(Clone)]
pub struct ConnectionManager {
client: Client,
connection: Arc<ArcSwap<SharedRedisFuture<MultiplexedConnection>>>,
runtime: Runtime,
retry_strategy: ExponentialBuilder,
connection_config: AsyncConnectionConfig,
}
type CloneableRedisResult<T> = Result<T, Arc<RedisError>>;
type SharedRedisFuture<T> = Shared<BoxFuture<'static, CloneableRedisResult<T>>>;
macro_rules! reconnect_if_dropped {
($self:expr, $result:expr, $current:expr) => {
if let Err(ref e) = $result {
if e.is_unrecoverable_error() {
$self.reconnect($current);
}
}
};
}
macro_rules! reconnect_if_io_error {
($self:expr, $result:expr, $current:expr) => {
if let Err(e) = $result {
if e.is_io_error() {
$self.reconnect($current);
}
return Err(e);
}
};
}
impl ConnectionManager {
pub async fn new(client: Client) -> RedisResult<Self> {
let config = ConnectionManagerConfig::new();
Self::new_with_config(client, config).await
}
#[deprecated(note = "Use `new_with_config`")]
pub async fn new_with_backoff(
client: Client,
exponent_base: u64,
factor: u64,
number_of_retries: usize,
) -> RedisResult<Self> {
let config = ConnectionManagerConfig::new()
.set_exponent_base(exponent_base)
.set_factor(factor)
.set_number_of_retries(number_of_retries);
Self::new_with_config(client, config).await
}
#[deprecated(note = "Use `new_with_config`")]
pub async fn new_with_backoff_and_timeouts(
client: Client,
exponent_base: u64,
factor: u64,
number_of_retries: usize,
response_timeout: std::time::Duration,
connection_timeout: std::time::Duration,
) -> RedisResult<Self> {
let config = ConnectionManagerConfig::new()
.set_exponent_base(exponent_base)
.set_factor(factor)
.set_number_of_retries(number_of_retries)
.set_response_timeout(response_timeout)
.set_connection_timeout(connection_timeout);
Self::new_with_config(client, config).await
}
pub async fn new_with_config(
client: Client,
config: ConnectionManagerConfig,
) -> RedisResult<Self> {
let runtime = Runtime::locate();
let mut retry_strategy = ExponentialBuilder::default()
.with_factor(config.factor as f32)
.with_max_times(config.number_of_retries)
.with_jitter();
if let Some(max_delay) = config.max_delay {
retry_strategy =
retry_strategy.with_max_delay(std::time::Duration::from_millis(max_delay));
}
let mut connection_config = AsyncConnectionConfig::new();
if let Some(connection_timeout) = config.connection_timeout {
connection_config = connection_config.set_connection_timeout(connection_timeout);
}
if let Some(response_timeout) = config.response_timeout {
connection_config = connection_config.set_response_timeout(response_timeout);
}
if let Some(push_sender) = config.push_sender.clone() {
check_resp3!(
client.connection_info.redis.protocol,
"Can only pass push sender to a connection using RESP3"
);
connection_config = connection_config.set_push_sender_internal(push_sender);
}
let connection =
Self::new_connection(client.clone(), retry_strategy, &connection_config).await?;
Ok(Self {
client,
connection: Arc::new(ArcSwap::from_pointee(
future::ok(connection).boxed().shared(),
)),
runtime,
retry_strategy,
connection_config,
})
}
async fn new_connection(
client: Client,
exponential_backoff: ExponentialBuilder,
connection_config: &AsyncConnectionConfig,
) -> RedisResult<MultiplexedConnection> {
let connection_config = connection_config.clone();
let get_conn = || async {
client
.get_multiplexed_async_connection_with_config(&connection_config)
.await
};
get_conn
.retry(exponential_backoff)
.sleep(|duration| async move { Runtime::locate().sleep(duration).await })
.await
}
fn reconnect(&self, current: arc_swap::Guard<Arc<SharedRedisFuture<MultiplexedConnection>>>) {
let client = self.client.clone();
let retry_strategy = self.retry_strategy;
let connection_config = self.connection_config.clone();
let new_connection: SharedRedisFuture<MultiplexedConnection> = async move {
let con = Self::new_connection(client, retry_strategy, &connection_config).await?;
Ok(con)
}
.boxed()
.shared();
let new_connection_arc = Arc::new(new_connection.clone());
let prev = self
.connection
.compare_and_swap(¤t, new_connection_arc);
if Arc::ptr_eq(&prev, ¤t) {
self.runtime.spawn(new_connection.map(|_| ()));
}
}
pub async fn send_packed_command(&mut self, cmd: &Cmd) -> RedisResult<Value> {
let guard = self.connection.load();
let connection_result = (**guard)
.clone()
.await
.map_err(|e| e.clone_mostly("Reconnecting failed"));
reconnect_if_io_error!(self, connection_result, guard);
let result = connection_result?.send_packed_command(cmd).await;
reconnect_if_dropped!(self, &result, guard);
result
}
pub async fn send_packed_commands(
&mut self,
cmd: &crate::Pipeline,
offset: usize,
count: usize,
) -> RedisResult<Vec<Value>> {
let guard = self.connection.load();
let connection_result = (**guard)
.clone()
.await
.map_err(|e| e.clone_mostly("Reconnecting failed"));
reconnect_if_io_error!(self, connection_result, guard);
let result = connection_result?
.send_packed_commands(cmd, offset, count)
.await;
reconnect_if_dropped!(self, &result, guard);
result
}
pub async fn subscribe(&mut self, channel_name: impl ToRedisArgs) -> RedisResult<()> {
check_resp3!(self.client.connection_info.redis.protocol);
let mut cmd = cmd("SUBSCRIBE");
cmd.arg(channel_name);
cmd.exec_async(self).await?;
Ok(())
}
pub async fn unsubscribe(&mut self, channel_name: impl ToRedisArgs) -> RedisResult<()> {
check_resp3!(self.client.connection_info.redis.protocol);
let mut cmd = cmd("UNSUBSCRIBE");
cmd.arg(channel_name);
cmd.exec_async(self).await?;
Ok(())
}
pub async fn psubscribe(&mut self, channel_pattern: impl ToRedisArgs) -> RedisResult<()> {
check_resp3!(self.client.connection_info.redis.protocol);
let mut cmd = cmd("PSUBSCRIBE");
cmd.arg(channel_pattern);
cmd.exec_async(self).await?;
Ok(())
}
pub async fn punsubscribe(&mut self, channel_pattern: impl ToRedisArgs) -> RedisResult<()> {
check_resp3!(self.client.connection_info.redis.protocol);
let mut cmd = cmd("PUNSUBSCRIBE");
cmd.arg(channel_pattern);
cmd.exec_async(self).await?;
Ok(())
}
}
impl ConnectionLike for ConnectionManager {
fn req_packed_command<'a>(&'a mut self, cmd: &'a Cmd) -> RedisFuture<'a, Value> {
(async move { self.send_packed_command(cmd).await }).boxed()
}
fn req_packed_commands<'a>(
&'a mut self,
cmd: &'a crate::Pipeline,
offset: usize,
count: usize,
) -> RedisFuture<'a, Vec<Value>> {
(async move { self.send_packed_commands(cmd, offset, count).await }).boxed()
}
fn get_db(&self) -> i64 {
self.client.connection_info().redis.db
}
}