use crate::refresh::{RefreshTokenError, RefreshTokenStore, TokenBlacklist};
use redis::aio::ConnectionManager;
use redis::AsyncCommands;
use std::fmt;
use std::time::Duration;
#[derive(Clone)]
pub struct RedisConfig {
pub url: String,
pub key_prefix_ver: String,
pub key_prefix_bl: String,
pub connection_timeout: Duration,
pub command_timeout: Duration,
}
impl Default for RedisConfig {
fn default() -> Self {
Self {
url: "redis://127.0.0.1:6379".to_string(),
key_prefix_ver: "sso:ver".to_string(),
key_prefix_bl: "sso:bl".to_string(),
connection_timeout: Duration::from_secs(3),
command_timeout: Duration::from_secs(2),
}
}
}
impl RedisConfig {
pub fn from_url(url: impl Into<String>) -> Self {
Self {
url: url.into(),
..Default::default()
}
}
fn ver_key(&self, user_id: i64) -> String {
format!("{}:{}", self.key_prefix_ver, user_id)
}
fn bl_key(&self, jti: &str) -> String {
format!("{}:{}", self.key_prefix_bl, jti)
}
}
impl fmt::Debug for RedisConfig {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let redacted_url = redact_redis_url(&self.url);
f.debug_struct("RedisConfig")
.field("url", &redacted_url)
.field("key_prefix_ver", &self.key_prefix_ver)
.field("key_prefix_bl", &self.key_prefix_bl)
.field("connection_timeout", &self.connection_timeout)
.field("command_timeout", &self.command_timeout)
.finish()
}
}
fn redact_redis_url(url: &str) -> String {
if let Some(at_pos) = url.find('@') {
if let Some(scheme_end) = url.find("://") {
let password_start = scheme_end + 3;
if at_pos > password_start {
let (before, after) = url.split_at(at_pos);
let scheme = &before[..password_start];
return format!("{}[REDACTED]{}", scheme, after);
}
}
}
url.to_string()
}
pub struct RedisRefreshTokenStore {
conn: ConnectionManager,
config: RedisConfig,
}
impl RedisRefreshTokenStore {
pub async fn new(config: RedisConfig) -> Result<Self, RefreshTokenError> {
let client = redis::Client::open(config.url.as_str())
.map_err(|e| RefreshTokenError::Cache(format!("redis client open failed: {e}")))?;
let conn = tokio::time::timeout(config.connection_timeout, client.get_connection_manager())
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis connect failed: {e}")))?;
Ok(Self { conn, config })
}
}
#[async_trait::async_trait]
impl RefreshTokenStore for RedisRefreshTokenStore {
async fn get_version(&self, user_id: i64) -> Result<u64, RefreshTokenError> {
let key = self.config.ver_key(user_id);
let mut conn = self.conn.clone();
let result: Option<u64> = tokio::time::timeout(
self.config.command_timeout,
conn.get::<&str, Option<u64>>(&key),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis GET failed: {e}")))?;
Ok(result.unwrap_or(0))
}
async fn increment_version(&self, user_id: i64) -> Result<u64, RefreshTokenError> {
let key = self.config.ver_key(user_id);
let mut conn = self.conn.clone();
let new_version: u64 = tokio::time::timeout(
self.config.command_timeout,
conn.incr::<&str, u64, u64>(&key, 1),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis INCR failed: {e}")))?;
Ok(new_version)
}
}
pub struct RedisTokenBlacklist {
conn: ConnectionManager,
config: RedisConfig,
}
impl RedisTokenBlacklist {
pub async fn new(config: RedisConfig) -> Result<Self, RefreshTokenError> {
let client = redis::Client::open(config.url.as_str())
.map_err(|e| RefreshTokenError::Cache(format!("redis client open failed: {e}")))?;
let conn = tokio::time::timeout(config.connection_timeout, client.get_connection_manager())
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis connect failed: {e}")))?;
Ok(Self { conn, config })
}
}
#[async_trait::async_trait]
impl TokenBlacklist for RedisTokenBlacklist {
async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RefreshTokenError> {
if ttl_secs == 0 {
return Ok(());
}
let key = self.config.bl_key(jti);
let mut conn = self.conn.clone();
tokio::time::timeout(
self.config.command_timeout,
conn.set_ex::<&str, &str, ()>(&key, "1", ttl_secs),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis SETEX failed: {e}")))?;
Ok(())
}
async fn is_revoked(&self, jti: &str) -> Result<bool, RefreshTokenError> {
let key = self.config.bl_key(jti);
let mut conn = self.conn.clone();
let exists: bool =
tokio::time::timeout(self.config.command_timeout, conn.exists::<&str, bool>(&key))
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis EXISTS failed: {e}")))?;
Ok(exists)
}
}
pub async fn create_redis_stores(
config: RedisConfig,
) -> Result<
(
std::sync::Arc<dyn RefreshTokenStore>,
std::sync::Arc<dyn TokenBlacklist>,
),
RefreshTokenError,
> {
let store = RedisRefreshTokenStore::new(config.clone()).await?;
let blacklist = RedisTokenBlacklist::new(config).await?;
Ok((std::sync::Arc::new(store), std::sync::Arc::new(blacklist)))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_redis_config_default() {
let config = RedisConfig::default();
assert_eq!(config.url, "redis://127.0.0.1:6379");
assert_eq!(config.key_prefix_ver, "sso:ver");
assert_eq!(config.key_prefix_bl, "sso:bl");
assert_eq!(config.connection_timeout, Duration::from_secs(3));
assert_eq!(config.command_timeout, Duration::from_secs(2));
}
#[test]
fn test_redis_config_from_url() {
let config = RedisConfig::from_url("redis://localhost:6380/1");
assert_eq!(config.url, "redis://localhost:6380/1");
assert_eq!(config.key_prefix_ver, "sso:ver");
}
#[test]
fn test_redis_config_debug_redacts_password() {
let config = RedisConfig::from_url("redis://:secret_pass@127.0.0.1:6379");
let debug_str = format!("{:?}", config);
assert!(debug_str.contains("[REDACTED]"));
assert!(!debug_str.contains("secret_pass"));
}
#[test]
fn test_redis_config_debug_no_password() {
let config = RedisConfig::from_url("redis://127.0.0.1:6379");
let debug_str = format!("{:?}", config);
assert!(!debug_str.contains("[REDACTED]"));
assert!(debug_str.contains("127.0.0.1:6379"));
}
#[test]
fn test_ver_key_format() {
let config = RedisConfig::default();
assert_eq!(config.ver_key(1), "sso:ver:1");
assert_eq!(config.ver_key(42), "sso:ver:42");
}
#[test]
fn test_bl_key_format() {
let config = RedisConfig::default();
assert_eq!(config.bl_key("abc123"), "sso:bl:abc123");
}
#[test]
fn test_redact_redis_url_with_password() {
let redacted = redact_redis_url("redis://:mypassword@host:6379/0");
assert!(redacted.contains("[REDACTED]"));
assert!(!redacted.contains("mypassword"));
assert!(redacted.contains("host:6379"));
}
#[test]
fn test_redact_redis_url_without_password() {
let redacted = redact_redis_url("redis://127.0.0.1:6379");
assert_eq!(redacted, "redis://127.0.0.1:6379");
}
#[test]
fn test_redact_redis_url_with_user_and_password() {
let redacted = redact_redis_url("redis://user:pass@host:6379");
assert!(redacted.contains("[REDACTED]"));
assert!(!redacted.contains("pass"));
}
}