feat(cache-redis): expose the pooled connection handling for reuse

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yujong Lee 2026-09-21 20:24:23 +00:00
parent 769917a7e8
commit 2b257bd9a3
3 changed files with 72 additions and 53 deletions

View file

@ -19,7 +19,7 @@ const DEFAULT_TTL: Duration = Duration::from_secs(600);
const REDIS_TIMEOUT: Duration = Duration::from_secs(5); const REDIS_TIMEOUT: Duration = Duration::from_secs(5);
const REDIS_POOL_SIZE: u32 = 16; const REDIS_POOL_SIZE: u32 = 16;
struct PooledConnection { pub struct PooledConnection {
connection: redis::Connection, connection: redis::Connection,
failed: bool, failed: bool,
} }
@ -27,16 +27,19 @@ struct PooledConnection {
/// Pools connections without a checkout PING, which would double every operation's round trips. /// Pools connections without a checkout PING, which would double every operation's round trips.
/// A timed-out command leaves its reply on the socket while redis still reports the connection /// A timed-out command leaves its reply on the socket while redis still reports the connection
/// open, so any connection whose operation failed is discarded instead of being reused. /// open, so any connection whose operation failed is discarded instead of being reused.
struct ConnectionManager(redis::Client); pub struct ConnectionManager {
client: redis::Client,
timeout: Duration,
}
impl r2d2::ManageConnection for ConnectionManager { impl r2d2::ManageConnection for ConnectionManager {
type Connection = PooledConnection; type Connection = PooledConnection;
type Error = redis::RedisError; type Error = redis::RedisError;
fn connect(&self) -> Result<PooledConnection, redis::RedisError> { fn connect(&self) -> Result<PooledConnection, redis::RedisError> {
let connection = self.0.get_connection()?; let connection = self.client.get_connection()?;
connection.set_read_timeout(Some(REDIS_TIMEOUT))?; connection.set_read_timeout(Some(self.timeout))?;
connection.set_write_timeout(Some(REDIS_TIMEOUT))?; connection.set_write_timeout(Some(self.timeout))?;
Ok(PooledConnection { Ok(PooledConnection {
connection, connection,
failed: false, failed: false,
@ -68,12 +71,12 @@ const CLAIM_SCRIPT: &str = concat!(
); );
const CLAIM_ATTEMPTS: usize = 8; const CLAIM_ATTEMPTS: usize = 8;
enum Connections<C> { pub enum Connections<C> {
Pool(r2d2::Pool<ConnectionManager>), Pool(r2d2::Pool<ConnectionManager>),
Fixed(Mutex<C>), Fixed(Mutex<C>),
} }
struct ConnectionRef<'a>(&'a mut dyn redis::ConnectionLike); pub struct ConnectionRef<'a>(&'a mut dyn redis::ConnectionLike);
impl redis::ConnectionLike for ConnectionRef<'_> { impl redis::ConnectionLike for ConnectionRef<'_> {
fn req_packed_command(&mut self, cmd: &[u8]) -> redis::RedisResult<redis::Value> { fn req_packed_command(&mut self, cmd: &[u8]) -> redis::RedisResult<redis::Value> {
@ -110,7 +113,23 @@ impl<C> Connections<C>
where where
C: redis::ConnectionLike + Send + 'static, C: redis::ConnectionLike + Send + 'static,
{ {
fn execute<T>( pub fn pooled(url: &str, timeout: Duration, pool_size: u32) -> Result<Self, Error> {
let client = redis::Client::open(url).map_err(|_| Error::Unavailable)?;
let pool = r2d2::Pool::builder()
.max_size(pool_size)
.min_idle(Some(0))
.connection_timeout(timeout)
.test_on_check_out(false)
.build(ConnectionManager { client, timeout })
.map_err(|_| Error::Unavailable)?;
Ok(Self::Pool(pool))
}
pub fn fixed(connection: C) -> Self {
Self::Fixed(Mutex::new(connection))
}
pub fn execute<T>(
&self, &self,
operation: impl FnOnce(&mut ConnectionRef<'_>) -> Result<T, Error>, operation: impl FnOnce(&mut ConnectionRef<'_>) -> Result<T, Error>,
) -> Result<T, Error> { ) -> Result<T, Error> {
@ -127,6 +146,16 @@ where
} }
} }
} }
pub async fn run_blocking<T, F>(connections: Arc<Self>, operation: F) -> Result<T, Error>
where
T: Send + 'static,
F: FnOnce(&mut ConnectionRef<'_>) -> Result<T, Error> + Send + 'static,
{
tokio::task::spawn_blocking(move || connections.execute(operation))
.await
.map_err(|_| Error::Unavailable)?
}
} }
pub struct RedisCache<S, C = redis::Connection> { pub struct RedisCache<S, C = redis::Connection> {
@ -138,16 +167,8 @@ pub struct RedisCache<S, C = redis::Connection> {
impl<S: CacheCodec> RedisCache<S> { impl<S: CacheCodec> RedisCache<S> {
pub fn new(url: &str, default_ttl: Option<Duration>, codec: S) -> Result<Self, Error> { pub fn new(url: &str, default_ttl: Option<Duration>, codec: S) -> Result<Self, Error> {
let client = redis::Client::open(url).map_err(|_| Error::Unavailable)?;
let pool = r2d2::Pool::builder()
.max_size(REDIS_POOL_SIZE)
.min_idle(Some(0))
.connection_timeout(REDIS_TIMEOUT)
.test_on_check_out(false)
.build(ConnectionManager(client))
.map_err(|_| Error::Unavailable)?;
Ok(Self { Ok(Self {
connections: Arc::new(Connections::Pool(pool)), connections: Arc::new(Connections::pooled(url, REDIS_TIMEOUT, REDIS_POOL_SIZE)?),
default_ttl: default_ttl.unwrap_or(DEFAULT_TTL), default_ttl: default_ttl.unwrap_or(DEFAULT_TTL),
codec, codec,
namespace: None, namespace: None,
@ -162,7 +183,7 @@ where
{ {
pub fn with_connection(connection: C, default_ttl: Option<Duration>, codec: S) -> Self { pub fn with_connection(connection: C, default_ttl: Option<Duration>, codec: S) -> Self {
Self { Self {
connections: Arc::new(Connections::Fixed(Mutex::new(connection))), connections: Arc::new(Connections::fixed(connection)),
default_ttl: default_ttl.unwrap_or(DEFAULT_TTL), default_ttl: default_ttl.unwrap_or(DEFAULT_TTL),
codec, codec,
namespace: None, namespace: None,
@ -241,20 +262,14 @@ where
} }
fn ttl_seconds(ttl: Duration) -> u64 { fn ttl_seconds(ttl: Duration) -> u64 {
ttl.as_secs() ttl_seconds(ttl)
.saturating_add(u64::from(ttl.subsec_nanos() > 0))
.max(1)
} }
}
async fn run_blocking<T, F>(connections: Arc<Connections<C>>, operation: F) -> Result<T, Error> pub fn ttl_seconds(ttl: Duration) -> u64 {
where ttl.as_secs()
T: Send + 'static, .saturating_add(u64::from(ttl.subsec_nanos() > 0))
F: FnOnce(&mut ConnectionRef<'_>) -> Result<T, Error> + Send + 'static, .max(1)
{
tokio::task::spawn_blocking(move || connections.execute(operation))
.await
.map_err(|_| Error::Unavailable)?
}
} }
fn namespaced_key(namespace: Option<&str>, key: &str) -> String { fn namespaced_key(namespace: Option<&str>, key: &str) -> String {
@ -313,7 +328,7 @@ where
let payload = self.codec.encode(&value)?; let payload = self.codec.encode(&value)?;
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl)); let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
connection connection
.set_ex::<_, _, ()>(key, payload, ttl) .set_ex::<_, _, ()>(key, payload, ttl)
.map_err(|_| Error::Unavailable) .map_err(|_| Error::Unavailable)
@ -327,7 +342,7 @@ where
_: &ExactCacheContext, _: &ExactCacheContext,
) -> Result<Option<Self::Value>, Error> { ) -> Result<Option<Self::Value>, Error> {
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
let value = Self::run_blocking(Arc::clone(&self.connections), move |connection| { let value = Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
connection connection
.get::<_, redis::Value>(key) .get::<_, redis::Value>(key)
.map_err(|_| Error::Unavailable) .map_err(|_| Error::Unavailable)
@ -350,7 +365,7 @@ where
}) })
.collect::<Result<Vec<_>, _>>()?; .collect::<Result<Vec<_>, _>>()?;
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl)); let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut pipeline = redis::pipe(); let mut pipeline = redis::pipe();
for (key, payload) in entries { for (key, payload) in entries {
pipeline pipeline
@ -372,7 +387,7 @@ where
} }
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> { async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
match Self::run_blocking(Arc::clone(&self.connections), |connection| { match Connections::run_blocking(Arc::clone(&self.connections), |connection| {
Ok(match redis::cmd("PING").query::<String>(connection) { Ok(match redis::cmd("PING").query::<String>(connection) {
Ok(_) => CacheConnectionResult { Ok(_) => CacheConnectionResult {
status: CacheConnectionStatus::Success, status: CacheConnectionStatus::Success,
@ -433,7 +448,7 @@ where
.iter() .iter()
.map(|key| self.namespaced_key(key)) .map(|key| self.namespaced_key(key))
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| { let values = Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("MGET") redis::cmd("MGET")
.arg(keys) .arg(keys)
.query::<Vec<redis::Value>>(connection) .query::<Vec<redis::Value>>(connection)
@ -460,7 +475,7 @@ where
async fn async_delete_cache(&self, key: &str) -> Result<(), Error> { async fn async_delete_cache(&self, key: &str) -> Result<(), Error> {
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
connection.del::<_, ()>(key).map_err(|_| Error::Unavailable) connection.del::<_, ()>(key).map_err(|_| Error::Unavailable)
}) })
.await .await
@ -480,7 +495,7 @@ where
async fn async_flush_cache(&self) -> Result<(), Error> { async fn async_flush_cache(&self) -> Result<(), Error> {
let pattern = self.namespaced_pattern()?; let pattern = self.namespaced_pattern()?;
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
Self::flush_matching(connection, &pattern) Self::flush_matching(connection, &pattern)
}) })
.await .await
@ -512,7 +527,7 @@ where
) -> Result<f64, Error> { ) -> Result<f64, Error> {
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl)); let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
increment(connection, key, amount, ttl) increment(connection, key, amount, ttl)
}) })
.await .await
@ -623,7 +638,7 @@ where
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl)); let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl));
let codec = self.codec.clone(); let codec = self.codec.clone();
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
claim(connection, &codec, &key, candidate, &eligible, ttl) claim(connection, &codec, &key, candidate, &eligible, ttl)
}) })
.await .await

View file

@ -144,7 +144,7 @@ where
.into_iter() .into_iter()
.map(|key| self.namespaced_key(&key)) .map(|key| self.namespaced_key(&key))
.collect::<Vec<_>>(); .collect::<Vec<_>>();
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
connection.del(keys).map_err(|_| Error::Unavailable) connection.del(keys).map_err(|_| Error::Unavailable)
}) })
.await .await
@ -172,7 +172,7 @@ where
.iter() .iter()
.map(|key| self.namespaced_key(key)) .map(|key| self.namespaced_key(key))
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| { let values = Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("MGET") redis::cmd("MGET")
.arg(keys) .arg(keys)
.query::<Vec<redis::Value>>(connection) .query::<Vec<redis::Value>>(connection)
@ -192,7 +192,7 @@ where
} }
pub async fn ping(&self) -> Result<bool, Error> { pub async fn ping(&self) -> Result<bool, Error> {
Self::run_blocking(Arc::clone(&self.connections), |connection| { Connections::run_blocking(Arc::clone(&self.connections), |connection| {
redis::cmd("PING") redis::cmd("PING")
.query::<String>(connection) .query::<String>(connection)
.map(|response| response == "PONG") .map(|response| response == "PONG")
@ -203,7 +203,7 @@ where
pub async fn async_get_ttl(&self, key: &str) -> Result<Option<i64>, Error> { pub async fn async_get_ttl(&self, key: &str) -> Result<Option<i64>, Error> {
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
let ttl = Self::run_blocking(Arc::clone(&self.connections), move |connection| { let ttl = Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("TTL") redis::cmd("TTL")
.arg(key) .arg(key)
.query::<i64>(connection) .query::<i64>(connection)
@ -215,7 +215,7 @@ where
pub async fn async_scan_iter(&self, pattern: &str, count: usize) -> Result<Vec<String>, Error> { pub async fn async_scan_iter(&self, pattern: &str, count: usize) -> Result<Vec<String>, Error> {
let pattern = format!("{}*", self.namespaced_key(pattern)); let pattern = format!("{}*", self.namespaced_key(pattern));
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut cursor = 0u64; let mut cursor = 0u64;
let mut matches = Vec::new(); let mut matches = Vec::new();
loop { loop {
@ -249,7 +249,7 @@ where
} }
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(ttl.unwrap_or(self.default_ttl)); let ttl = Self::ttl_seconds(ttl.unwrap_or(self.default_ttl));
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut pipeline = redis::pipe(); let mut pipeline = redis::pipe();
pipeline.cmd("SADD").arg(&key).arg(values); pipeline.cmd("SADD").arg(&key).arg(values);
pipeline.cmd("EXPIRE").arg(&key).arg(ttl).ignore(); pipeline.cmd("EXPIRE").arg(&key).arg(ttl).ignore();
@ -266,7 +266,7 @@ where
return Err(Error::InvalidEntry); return Err(Error::InvalidEntry);
} }
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("RPUSH") redis::cmd("RPUSH")
.arg(key) .arg(key)
.arg(values) .arg(values)
@ -292,7 +292,7 @@ where
if operations.is_empty() { if operations.is_empty() {
return Ok(Vec::new()); return Ok(Vec::new());
} }
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut pipeline = redis::pipe(); let mut pipeline = redis::pipe();
for (key, values) in operations { for (key, values) in operations {
pipeline.cmd("RPUSH").arg(key).arg(values); pipeline.cmd("RPUSH").arg(key).arg(values);
@ -309,7 +309,7 @@ where
) -> Result<RedisLpopResult, Error> { ) -> Result<RedisLpopResult, Error> {
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
let multiple = count.is_some(); let multiple = count.is_some();
let value = Self::run_blocking(Arc::clone(&self.connections), move |connection| { let value = Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut command = redis::cmd("LPOP"); let mut command = redis::cmd("LPOP");
command.arg(key); command.arg(key);
if let Some(count) = count { if let Some(count) = count {
@ -338,7 +338,7 @@ where
.iter() .iter()
.map(|(_, count)| count.is_some()) .map(|(_, count)| count.is_some())
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| { let values = Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut pipeline = redis::pipe(); let mut pipeline = redis::pipe();
for (key, count) in operations { for (key, count) in operations {
let command = pipeline.cmd("LPOP").arg(key); let command = pipeline.cmd("LPOP").arg(key);
@ -368,7 +368,7 @@ where
.into_iter() .into_iter()
.map(|key| self.namespaced_key(&key)) .map(|key| self.namespaced_key(&key))
.collect::<Vec<_>>(); .collect::<Vec<_>>();
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("EVAL") redis::cmd("EVAL")
.arg(script) .arg(script)
.arg(keys.len()) .arg(keys.len())
@ -440,7 +440,7 @@ where
if operations.is_empty() { if operations.is_empty() {
return Ok(Vec::new()); return Ok(Vec::new());
} }
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
let mut pipeline = redis::pipe(); let mut pipeline = redis::pipe();
for (key, amount, ttl) in operations { for (key, amount, ttl) in operations {
pipeline.cmd("INCRBYFLOAT").arg(&key).arg(amount); pipeline.cmd("INCRBYFLOAT").arg(&key).arg(amount);
@ -461,7 +461,7 @@ where
) -> Result<i64, Error> { ) -> Result<i64, Error> {
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(ttl); let ttl = Self::ttl_seconds(ttl);
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
increment_with_floor(connection, key, amount, ttl) increment_with_floor(connection, key, amount, ttl)
}) })
.await .await
@ -475,7 +475,7 @@ where
) -> Result<f64, Error> { ) -> Result<f64, Error> {
let key = self.namespaced_key(key); let key = self.namespaced_key(key);
let ttl = Self::ttl_seconds(ttl.unwrap_or(self.default_ttl)); let ttl = Self::ttl_seconds(ttl.unwrap_or(self.default_ttl));
Self::run_blocking(Arc::clone(&self.connections), move |connection| { Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
redis::cmd("EVAL") redis::cmd("EVAL")
.arg(SET_MAX_SCRIPT) .arg(SET_MAX_SCRIPT)
.arg(1) .arg(1)

View file

@ -1,6 +1,10 @@
mod cache; mod cache;
mod topology; mod topology;
pub mod connection {
pub use crate::cache::{ConnectionRef, Connections, ttl_seconds};
}
pub use cache::{ pub use cache::{
RedisArg, RedisCache, RedisLpopOperation, RedisLpopResult, RedisRpushOperation, RedisScript, RedisArg, RedisCache, RedisLpopOperation, RedisLpopResult, RedisRpushOperation, RedisScript,
}; };