use crate::refresh::{
DeviceInfo, DeviceSession, DeviceSessionStore, 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 key_prefix_sessions: 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(),
key_prefix_sessions: "sso:sessions".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)
}
fn sessions_key(&self, user_id: i64) -> String {
format!("{}:{}", self.key_prefix_sessions, user_id)
}
}
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("key_prefix_sessions", &self.key_prefix_sessions)
.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 struct RedisDeviceSessionStore {
conn: ConnectionManager,
config: RedisConfig,
}
impl RedisDeviceSessionStore {
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 })
}
pub fn from_conn(conn: ConnectionManager, config: RedisConfig) -> Self {
Self { conn, config }
}
}
#[async_trait::async_trait]
impl DeviceSessionStore for RedisDeviceSessionStore {
async fn register_session(
&self,
user_id: i64,
device_id: &str,
device_info: &DeviceInfo,
jti: &str,
access_jti: &str,
) -> Result<(), RefreshTokenError> {
let now = chrono::Utc::now().timestamp();
let session = DeviceSession {
device_id: device_id.to_string(),
device_info: device_info.clone(),
jti: jti.to_string(),
access_jti: access_jti.to_string(),
created_at: now,
last_active: now,
};
let key = self.config.sessions_key(user_id);
let value = serde_json::to_string(&session)
.map_err(|e| RefreshTokenError::Cache(format!("json serialize failed: {e}")))?;
let mut conn = self.conn.clone();
tokio::time::timeout(
self.config.command_timeout,
conn.hset::<&str, &str, &str, ()>(&key, device_id, &value),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HSET failed: {e}")))?;
Ok(())
}
async fn get_sessions(&self, user_id: i64) -> Result<Vec<DeviceSession>, RefreshTokenError> {
let key = self.config.sessions_key(user_id);
let mut conn = self.conn.clone();
let map: std::collections::HashMap<String, String> = tokio::time::timeout(
self.config.command_timeout,
conn.hgetall::<&str, std::collections::HashMap<String, String>>(&key),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HGETALL failed: {e}")))?;
let mut sessions = Vec::with_capacity(map.len());
for (_, v) in map {
let session: DeviceSession = serde_json::from_str(&v)
.map_err(|e| RefreshTokenError::Cache(format!("json deserialize failed: {e}")))?;
sessions.push(session);
}
Ok(sessions)
}
async fn get_session(
&self,
user_id: i64,
device_id: &str,
) -> Result<Option<DeviceSession>, RefreshTokenError> {
let key = self.config.sessions_key(user_id);
let mut conn = self.conn.clone();
let value: Option<String> = tokio::time::timeout(
self.config.command_timeout,
conn.hget::<&str, &str, Option<String>>(&key, device_id),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HGET failed: {e}")))?;
match value {
Some(v) => {
let session: DeviceSession = serde_json::from_str(&v).map_err(|e| {
RefreshTokenError::Cache(format!("json deserialize failed: {e}"))
})?;
Ok(Some(session))
}
None => Ok(None),
}
}
async fn revoke_session(
&self,
user_id: i64,
device_id: &str,
) -> Result<Option<(String, String)>, RefreshTokenError> {
let key = self.config.sessions_key(user_id);
let mut conn = self.conn.clone();
let value: Option<String> = tokio::time::timeout(
self.config.command_timeout,
conn.hget::<&str, &str, Option<String>>(&key, device_id),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HGET failed: {e}")))?;
match value {
Some(v) => {
let session: DeviceSession = serde_json::from_str(&v).map_err(|e| {
RefreshTokenError::Cache(format!("json deserialize failed: {e}"))
})?;
tokio::time::timeout(
self.config.command_timeout,
conn.hdel::<&str, &str, ()>(&key, device_id),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HDEL failed: {e}")))?;
Ok(Some((session.jti, session.access_jti)))
}
None => Ok(None),
}
}
async fn update_last_active(
&self,
user_id: i64,
device_id: &str,
) -> Result<(), RefreshTokenError> {
let key = self.config.sessions_key(user_id);
let mut conn = self.conn.clone();
let value: Option<String> = tokio::time::timeout(
self.config.command_timeout,
conn.hget::<&str, &str, Option<String>>(&key, device_id),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HGET failed: {e}")))?;
match value {
Some(v) => {
let mut session: DeviceSession = serde_json::from_str(&v).map_err(|e| {
RefreshTokenError::Cache(format!("json deserialize failed: {e}"))
})?;
session.last_active = chrono::Utc::now().timestamp();
let new_value = serde_json::to_string(&session)
.map_err(|e| RefreshTokenError::Cache(format!("json serialize failed: {e}")))?;
tokio::time::timeout(
self.config.command_timeout,
conn.hset::<&str, &str, &str, ()>(&key, device_id, &new_value),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HSET failed: {e}")))?;
Ok(())
}
None => Ok(()),
}
}
async fn update_session_jti(
&self,
user_id: i64,
device_id: &str,
new_jti: &str,
) -> Result<(), RefreshTokenError> {
let key = self.config.sessions_key(user_id);
let mut conn = self.conn.clone();
let value: Option<String> = tokio::time::timeout(
self.config.command_timeout,
conn.hget::<&str, &str, Option<String>>(&key, device_id),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HGET failed: {e}")))?;
match value {
Some(v) => {
let mut session: DeviceSession = serde_json::from_str(&v).map_err(|e| {
RefreshTokenError::Cache(format!("json deserialize failed: {e}"))
})?;
session.jti = new_jti.to_string();
session.last_active = chrono::Utc::now().timestamp();
let new_value = serde_json::to_string(&session)
.map_err(|e| RefreshTokenError::Cache(format!("json serialize failed: {e}")))?;
tokio::time::timeout(
self.config.command_timeout,
conn.hset::<&str, &str, &str, ()>(&key, device_id, &new_value),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HSET failed: {e}")))?;
Ok(())
}
None => Ok(()),
}
}
async fn cleanup_expired(
&self,
user_id: i64,
ttl_secs: i64,
) -> Result<Vec<(String, String)>, RefreshTokenError> {
let key = self.config.sessions_key(user_id);
let mut conn = self.conn.clone();
let map: std::collections::HashMap<String, String> = tokio::time::timeout(
self.config.command_timeout,
conn.hgetall::<&str, std::collections::HashMap<String, String>>(&key),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HGETALL failed: {e}")))?;
let now = chrono::Utc::now().timestamp();
let mut expired_fields = Vec::new();
let mut jti_list = Vec::new();
for (field, v) in map {
let session: DeviceSession = serde_json::from_str(&v)
.map_err(|e| RefreshTokenError::Cache(format!("json deserialize failed: {e}")))?;
if session.last_active + ttl_secs < now {
jti_list.push((session.jti.clone(), session.access_jti.clone()));
expired_fields.push(field);
}
}
if !expired_fields.is_empty() {
for field in &expired_fields {
tokio::time::timeout(
self.config.command_timeout,
conn.hdel::<&str, &str, ()>(&key, field),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HDEL failed: {e}")))?;
}
tracing::debug!(
user_id,
count = expired_fields.len(),
"expired sessions cleaned"
);
}
Ok(jti_list)
}
async fn clear_user_sessions(
&self,
user_id: i64,
) -> Result<Vec<(String, String)>, RefreshTokenError> {
let key = self.config.sessions_key(user_id);
let mut conn = self.conn.clone();
let map: std::collections::HashMap<String, String> = tokio::time::timeout(
self.config.command_timeout,
conn.hgetall::<&str, std::collections::HashMap<String, String>>(&key),
)
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis HGETALL failed: {e}")))?;
let mut jti_list = Vec::with_capacity(map.len());
for (_, v) in map {
let session: DeviceSession = serde_json::from_str(&v)
.map_err(|e| RefreshTokenError::Cache(format!("json deserialize failed: {e}")))?;
jti_list.push((session.jti, session.access_jti));
}
tokio::time::timeout(self.config.command_timeout, conn.del::<&str, ()>(&key))
.await
.map_err(|_| RefreshTokenError::ServiceUnavailable)?
.map_err(|e| RefreshTokenError::Cache(format!("redis DEL failed: {e}")))?;
Ok(jti_list)
}
}
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)))
}
pub async fn create_redis_stores_with_devices(
config: RedisConfig,
) -> Result<
(
std::sync::Arc<dyn RefreshTokenStore>,
std::sync::Arc<dyn TokenBlacklist>,
std::sync::Arc<dyn DeviceSessionStore>,
),
RefreshTokenError,
> {
let store = RedisRefreshTokenStore::new(config.clone()).await?;
let blacklist = RedisTokenBlacklist::new(config.clone()).await?;
let device_store = RedisDeviceSessionStore::new(config).await?;
Ok((
std::sync::Arc::new(store),
std::sync::Arc::new(blacklist),
std::sync::Arc::new(device_store),
))
}
#[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"));
}
}