Skip to main content

sa_token_storage_redis/
lib.rs

1// Author: 金书记
2//
3//! # sa-token-storage-redis
4//!
5//! Redis存储实现
6//!
7//! 适用于:
8//! - 分布式部署
9//! - 需要数据持久化
10//! - 高性能要求的场景
11//!
12//! ## 使用方式
13//!
14//! ### 方式 1: 使用 Redis URL
15//! ```rust,ignore
16//! use sa_token_storage_redis::RedisStorage;
17//!
18//! // 无密码
19//! let storage = RedisStorage::new("redis://localhost:6379/0", "sa-token:").await?;
20//!
21//! // 有密码
22//! let storage = RedisStorage::new("redis://:password@localhost:6379/0", "sa-token:").await?;
23//! ```
24//!
25//! ### 方式 2: 使用配置结构体
26//! ```rust,ignore
27//! use sa_token_storage_redis::{RedisStorage, RedisConfig};
28//!
29//! let config = RedisConfig {
30//!     host: "localhost".to_string(),
31//!     port: 6379,
32//!     password: Some("your-password".to_string()),
33//!     database: 0,
34//!     pool_size: 10,
35//! };
36//!
37//! let storage = RedisStorage::from_config(config, "sa-token:").await?;
38//! ```
39
40use async_trait::async_trait;
41use redis::{AsyncCommands, Client, aio::ConnectionManager};
42use sa_token_adapter::storage::{SaStorage, ScanPage, StorageError, StorageResult};
43use serde::{Deserialize, Serialize};
44use std::time::Duration;
45
46/// Redis 配置
47#[derive(Debug, Clone, Serialize, Deserialize)]
48pub struct RedisConfig {
49    /// Redis 主机地址
50    #[serde(default = "default_host")]
51    pub host: String,
52
53    /// Redis 端口
54    #[serde(default = "default_port")]
55    pub port: u16,
56
57    /// Redis 密码(可选)
58    #[serde(default)]
59    pub password: Option<String>,
60
61    /// 数据库编号
62    #[serde(default)]
63    pub database: u8,
64
65    /// 连接池大小(暂未使用,保留用于未来扩展)
66    #[serde(default = "default_pool_size")]
67    pub pool_size: u32,
68}
69
70impl Default for RedisConfig {
71    fn default() -> Self {
72        Self {
73            host: default_host(),
74            port: default_port(),
75            password: None,
76            database: 0,
77            pool_size: default_pool_size(),
78        }
79    }
80}
81
82impl RedisConfig {
83    /// 转换为 Redis URL
84    ///
85    /// 支持的格式:
86    /// - `redis://localhost:6379/0` (无密码)
87    /// - `redis://:password@localhost:6379/0` (有密码)
88    pub fn to_url(&self) -> String {
89        if let Some(password) = &self.password {
90            format!(
91                "redis://:{}@{}:{}/{}",
92                password, self.host, self.port, self.database
93            )
94        } else {
95            format!("redis://{}:{}/{}", self.host, self.port, self.database)
96        }
97    }
98}
99
100fn default_host() -> String {
101    "localhost".to_string()
102}
103
104fn default_port() -> u16 {
105    6379
106}
107
108fn default_pool_size() -> u32 {
109    10
110}
111
112/// Redis存储实现
113#[derive(Clone)]
114pub struct RedisStorage {
115    client: ConnectionManager,
116    key_prefix: String,
117}
118
119impl std::fmt::Debug for RedisStorage {
120    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
121        f.debug_struct("RedisStorage")
122            .field("key_prefix", &self.key_prefix)
123            .finish_non_exhaustive()
124    }
125}
126
127impl RedisStorage {
128    /// 使用 Redis URL 创建存储
129    ///
130    /// # 参数
131    /// * `redis_url` - Redis 连接 URL
132    /// * `key_prefix` - 键前缀(例如:`sa-token:`)
133    ///
134    /// # URL 格式
135    /// - 无密码:`redis://localhost:6379/0`
136    /// - 有密码:`redis://:password@localhost:6379/0`
137    /// - 复杂密码:`redis://:Aq23-hjPwFB3mBDNFp3W1@localhost:6379/0`
138    ///
139    /// # 示例
140    /// ```rust,ignore
141    /// use sa_token_storage_redis::RedisStorage;
142    ///
143    /// // 无密码
144    /// let storage = RedisStorage::new("redis://localhost:6379/0", "sa-token:").await?;
145    ///
146    /// // 有密码
147    /// let storage = RedisStorage::new(
148    ///     "redis://:Aq23-hjPwFB3mBDNFp3W1@localhost:6379/0",
149    ///     "sa-token:"
150    /// ).await?;
151    /// ```
152    pub async fn new(redis_url: &str, key_prefix: impl Into<String>) -> StorageResult<Self> {
153        let client =
154            Client::open(redis_url).map_err(|e| StorageError::ConnectionError(e.to_string()))?;
155
156        let connection_manager = ConnectionManager::new(client)
157            .await
158            .map_err(|e| StorageError::ConnectionError(e.to_string()))?;
159
160        Ok(Self {
161            client: connection_manager,
162            key_prefix: key_prefix.into(),
163        })
164    }
165
166    /// 使用配置结构体创建存储
167    ///
168    /// # 参数
169    /// * `config` - Redis 配置
170    /// * `key_prefix` - 键前缀(例如:`sa-token:`)
171    ///
172    /// # 示例
173    /// ```rust,ignore
174    /// use sa_token_storage_redis::{RedisStorage, RedisConfig};
175    ///
176    /// let config = RedisConfig {
177    ///     host: "localhost".to_string(),
178    ///     port: 6379,
179    ///     password: Some("Aq23-hjPwFB3mBDNFp3W1".to_string()),
180    ///     database: 0,
181    ///     pool_size: 10,
182    /// };
183    ///
184    /// let storage = RedisStorage::from_config(config, "sa-token:").await?;
185    /// ```
186    pub async fn from_config(
187        config: RedisConfig,
188        key_prefix: impl Into<String>,
189    ) -> StorageResult<Self> {
190        let redis_url = config.to_url();
191        Self::new(&redis_url, key_prefix).await
192    }
193
194    /// 使用构建器模式创建存储
195    ///
196    /// # 示例
197    /// ```rust,ignore
198    /// use sa_token_storage_redis::RedisStorage;
199    ///
200    /// let storage = RedisStorage::builder()
201    ///     .host("localhost")
202    ///     .port(6379)
203    ///     .password("Aq23-hjPwFB3mBDNFp3W1")
204    ///     .database(0)
205    ///     .key_prefix("sa-token:")
206    ///     .build()
207    ///     .await?;
208    /// ```
209    pub fn builder() -> RedisStorageBuilder {
210        RedisStorageBuilder::default()
211    }
212
213    /// 获取完整的键名(带前缀)
214    fn full_key(&self, key: &str) -> String {
215        format!("{}{}", self.key_prefix, key)
216    }
217
218    /// 列表键使用独立前缀,避免与字符串键类型冲突
219    fn list_key(&self, key: &str) -> String {
220        format!("{}list:{}", self.key_prefix, key)
221    }
222
223    /// 将物理键剥离为逻辑键;前缀不匹配时返回 `None`
224    fn strip_prefix<'a>(&self, raw: &'a str) -> Option<&'a str> {
225        raw.strip_prefix(&self.key_prefix)
226    }
227
228    /// 便捷构造:物理前缀默认为空(逻辑键由 SaKeys 提供)
229    pub async fn connect(redis_url: &str) -> StorageResult<Self> {
230        Self::new(redis_url, "").await
231    }
232}
233
234/// Redis 存储构建器
235#[derive(Default)]
236pub struct RedisStorageBuilder {
237    config: RedisConfig,
238    key_prefix: Option<String>,
239}
240
241impl std::fmt::Debug for RedisStorageBuilder {
242    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
243        f.debug_struct("RedisStorageBuilder")
244            .field("key_prefix", &self.key_prefix)
245            .finish_non_exhaustive()
246    }
247}
248
249impl RedisStorageBuilder {
250    /// 设置 Redis 主机地址
251    pub fn host(mut self, host: impl Into<String>) -> Self {
252        self.config.host = host.into();
253        self
254    }
255
256    /// 设置 Redis 端口
257    pub fn port(mut self, port: u16) -> Self {
258        self.config.port = port;
259        self
260    }
261
262    /// 设置 Redis 密码
263    pub fn password(mut self, password: impl Into<String>) -> Self {
264        self.config.password = Some(password.into());
265        self
266    }
267
268    /// 设置数据库编号
269    pub fn database(mut self, database: u8) -> Self {
270        self.config.database = database;
271        self
272    }
273
274    /// 设置连接池大小(保留用于未来扩展)
275    pub fn pool_size(mut self, size: u32) -> Self {
276        self.config.pool_size = size;
277        self
278    }
279
280    /// 设置键前缀
281    pub fn key_prefix(mut self, prefix: impl Into<String>) -> Self {
282        self.key_prefix = Some(prefix.into());
283        self
284    }
285
286    /// 构建 RedisStorage(未设置 `key_prefix` 时默认为空字符串)
287    pub async fn build(self) -> StorageResult<RedisStorage> {
288        let key_prefix = self.key_prefix.unwrap_or_default();
289        RedisStorage::from_config(self.config, key_prefix).await
290    }
291}
292
293#[async_trait]
294impl SaStorage for RedisStorage {
295    async fn get(&self, key: &str) -> StorageResult<Option<String>> {
296        let mut conn = self.client.clone();
297        let full_key = self.full_key(key);
298
299        conn.get(&full_key)
300            .await
301            .map_err(|e| StorageError::OperationFailed(e.to_string()))
302    }
303
304    async fn set(&self, key: &str, value: &str, ttl: Option<Duration>) -> StorageResult<()> {
305        let mut conn = self.client.clone();
306        let full_key = self.full_key(key);
307
308        if let Some(ttl) = ttl {
309            conn.set_ex(&full_key, value, ttl.as_secs())
310                .await
311                .map_err(|e| StorageError::OperationFailed(e.to_string()))
312        } else {
313            conn.set(&full_key, value)
314                .await
315                .map_err(|e| StorageError::OperationFailed(e.to_string()))
316        }
317    }
318
319    async fn delete(&self, key: &str) -> StorageResult<()> {
320        let mut conn = self.client.clone();
321        let full_key = self.full_key(key);
322
323        conn.del(&full_key)
324            .await
325            .map_err(|e| StorageError::OperationFailed(e.to_string()))
326    }
327
328    async fn exists(&self, key: &str) -> StorageResult<bool> {
329        let mut conn = self.client.clone();
330        let full_key = self.full_key(key);
331
332        conn.exists(&full_key)
333            .await
334            .map_err(|e| StorageError::OperationFailed(e.to_string()))
335    }
336
337    async fn expire(&self, key: &str, ttl: Duration) -> StorageResult<()> {
338        let mut conn = self.client.clone();
339        let full_key = self.full_key(key);
340
341        conn.expire(&full_key, ttl.as_secs() as i64)
342            .await
343            .map_err(|e| StorageError::OperationFailed(e.to_string()))
344    }
345
346    async fn ttl(&self, key: &str) -> StorageResult<Option<Duration>> {
347        let mut conn = self.client.clone();
348        let full_key = self.full_key(key);
349
350        let ttl_secs: i64 = conn
351            .ttl(&full_key)
352            .await
353            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
354
355        match ttl_secs {
356            -2 => Ok(None), // 键不存在
357            -1 => Ok(None), // 永不过期
358            secs if secs > 0 => Ok(Some(Duration::from_secs(secs as u64))),
359            _ => Ok(Some(Duration::from_secs(0))),
360        }
361    }
362
363    async fn mget(&self, keys: &[&str]) -> StorageResult<Vec<Option<String>>> {
364        let mut conn = self.client.clone();
365        let full_keys: Vec<String> = keys.iter().map(|k| self.full_key(k)).collect();
366
367        // redis 1.x 的 `get` 只接受 ToSingleRedisArg,批量取值需走 `mget`
368        // redis 1.x's `get` only accepts ToSingleRedisArg; use `mget` for multi-key reads
369        conn.mget(&full_keys)
370            .await
371            .map_err(|e| StorageError::OperationFailed(e.to_string()))
372    }
373
374    async fn mset(&self, items: &[(&str, &str)], ttl: Option<Duration>) -> StorageResult<()> {
375        let mut conn = self.client.clone();
376        let full_items: Vec<(String, &str)> =
377            items.iter().map(|(k, v)| (self.full_key(k), *v)).collect();
378
379        // 使用 pipeline 批量操作
380        let mut pipe = redis::pipe();
381        for (key, value) in &full_items {
382            if let Some(ttl) = ttl {
383                pipe.set_ex(key, *value, ttl.as_secs());
384            } else {
385                pipe.set(key, *value);
386            }
387        }
388
389        pipe.query_async(&mut conn)
390            .await
391            .map_err(|e| StorageError::OperationFailed(e.to_string()))
392    }
393
394    async fn mdel(&self, keys: &[&str]) -> StorageResult<()> {
395        let mut conn = self.client.clone();
396        let full_keys: Vec<String> = keys.iter().map(|k| self.full_key(k)).collect();
397
398        conn.del(&full_keys)
399            .await
400            .map_err(|e| StorageError::OperationFailed(e.to_string()))
401    }
402
403    async fn incr(&self, key: &str) -> StorageResult<i64> {
404        let mut conn = self.client.clone();
405        let full_key = self.full_key(key);
406
407        conn.incr(&full_key, 1)
408            .await
409            .map_err(|e| StorageError::OperationFailed(e.to_string()))
410    }
411
412    async fn decr(&self, key: &str) -> StorageResult<i64> {
413        let mut conn = self.client.clone();
414        let full_key = self.full_key(key);
415
416        conn.decr(&full_key, 1)
417            .await
418            .map_err(|e| StorageError::OperationFailed(e.to_string()))
419    }
420
421    async fn clear(&self) -> StorageResult<()> {
422        let mut conn = self.client.clone();
423        let pattern = format!("{}*", self.key_prefix);
424        let mut cursor: u64 = 0;
425
426        loop {
427            let (next, keys): (u64, Vec<String>) = redis::cmd("SCAN")
428                .arg(cursor)
429                .arg("MATCH")
430                .arg(&pattern)
431                .arg("COUNT")
432                .arg(100)
433                .query_async(&mut conn)
434                .await
435                .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
436
437            if !keys.is_empty() {
438                conn.del::<_, ()>(&keys)
439                    .await
440                    .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
441            }
442
443            if next == 0 {
444                break;
445            }
446            cursor = next;
447        }
448
449        Ok(())
450    }
451
452    async fn set_if_absent(
453        &self,
454        key: &str,
455        value: &str,
456        ttl: Option<Duration>,
457    ) -> StorageResult<bool> {
458        let mut conn = self.client.clone();
459        let full_key = self.full_key(key);
460
461        let inserted: bool = if let Some(ttl) = ttl {
462            redis::cmd("SET")
463                .arg(&full_key)
464                .arg(value)
465                .arg("NX")
466                .arg("EX")
467                .arg(ttl.as_secs())
468                .query_async(&mut conn)
469                .await
470                .map_err(|e| StorageError::OperationFailed(e.to_string()))?
471        } else {
472            conn.set_nx(&full_key, value)
473                .await
474                .map_err(|e| StorageError::OperationFailed(e.to_string()))?
475        };
476
477        Ok(inserted)
478    }
479
480    async fn get_del(&self, key: &str) -> StorageResult<Option<String>> {
481        let mut conn = self.client.clone();
482        let full_key = self.full_key(key);
483
484        let value: Option<String> = redis::cmd("GETDEL")
485            .arg(&full_key)
486            .query_async(&mut conn)
487            .await
488            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
489
490        Ok(value)
491    }
492
493    async fn compare_and_swap(
494        &self,
495        key: &str,
496        expected: Option<&str>,
497        new_value: &str,
498        ttl: Option<Duration>,
499    ) -> StorageResult<bool> {
500        let mut conn = self.client.clone();
501        let full_key = self.full_key(key);
502        let expected_str = expected.unwrap_or("");
503        let ttl_secs = ttl.map(|d| d.as_secs()).unwrap_or(0);
504
505        // Lua 保证 GET + 比较 + SET 单键原子,避免 WATCH 竞态。
506        // expected=None 时 ARGV[1] 为空串:键不存在(GET 返回 false)视为匹配。
507        let script = r#"
508            local current = redis.call('GET', KEYS[1])
509            if current == false then current = '' end
510            if current ~= ARGV[1] then return 0 end
511            if tonumber(ARGV[3]) > 0 then
512                redis.call('SET', KEYS[1], ARGV[2], 'EX', ARGV[3])
513            else
514                redis.call('SET', KEYS[1], ARGV[2])
515            end
516            return 1
517        "#;
518
519        let swapped: i32 = redis::Script::new(script)
520            .key(&full_key)
521            .arg(expected_str)
522            .arg(new_value)
523            .arg(ttl_secs)
524            .invoke_async(&mut conn)
525            .await
526            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
527
528        Ok(swapped == 1)
529    }
530
531    async fn compare_and_delete(&self, key: &str, expected: &str) -> StorageResult<bool> {
532        let mut conn = self.client.clone();
533        let full_key = self.full_key(key);
534
535        let script = r#"
536            local current = redis.call('GET', KEYS[1])
537            if current == false then current = '' end
538            if current ~= ARGV[1] then return 0 end
539            redis.call('DEL', KEYS[1])
540            return 1
541        "#;
542
543        let deleted: i32 = redis::Script::new(script)
544            .key(&full_key)
545            .arg(expected)
546            .invoke_async(&mut conn)
547            .await
548            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
549
550        Ok(deleted == 1)
551    }
552
553    async fn list_push(
554        &self,
555        key: &str,
556        member: &str,
557        unique: bool,
558        ttl: Option<Duration>,
559    ) -> StorageResult<usize> {
560        // 【A1-2】Lua 原子去重:LRANGE + 判断 + RPUSH + EXPIRE 在同一脚本内完成
561        let mut conn = self.client.clone();
562        let list_key = self.list_key(key);
563        let ttl_secs = ttl.map(|d| d.as_secs()).unwrap_or(0);
564
565        let script = r#"
566            local list_key = KEYS[1]
567            local member = ARGV[1]
568            local unique = tonumber(ARGV[2])
569            local ttl_secs = tonumber(ARGV[3])
570
571            if unique == 1 then
572                local items = redis.call('LRANGE', list_key, 0, -1)
573                for i, v in ipairs(items) do
574                    if v == member then
575                        if ttl_secs > 0 then
576                            redis.call('EXPIRE', list_key, ttl_secs)
577                        end
578                        return #items
579                    end
580                end
581            end
582
583            local len = redis.call('RPUSH', list_key, member)
584            if ttl_secs > 0 then
585                redis.call('EXPIRE', list_key, ttl_secs)
586            end
587            return len
588        "#;
589
590        let len: usize = redis::Script::new(script)
591            .key(&list_key)
592            .arg(member)
593            .arg(if unique { 1 } else { 0 })
594            .arg(ttl_secs)
595            .invoke_async(&mut conn)
596            .await
597            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
598
599        Ok(len)
600    }
601
602    async fn list_remove(&self, key: &str, member: &str) -> StorageResult<bool> {
603        let mut conn = self.client.clone();
604        let list_key = self.list_key(key);
605
606        let removed: i64 = conn
607            .lrem(&list_key, 0, member)
608            .await
609            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
610
611        Ok(removed > 0)
612    }
613
614    async fn list_range(
615        &self,
616        key: &str,
617        start: usize,
618        limit: Option<usize>,
619    ) -> StorageResult<Vec<String>> {
620        let mut conn = self.client.clone();
621        let list_key = self.list_key(key);
622        let stop = match limit {
623            Some(l) => {
624                let end = start.saturating_add(l).saturating_sub(1);
625                isize::try_from(end.min(isize::MAX as usize)).unwrap_or(isize::MAX)
626            }
627            None => -1,
628        };
629
630        conn.lrange(&list_key, start as isize, stop)
631            .await
632            .map_err(|e| StorageError::OperationFailed(e.to_string()))
633    }
634
635    async fn list_len(&self, key: &str) -> StorageResult<usize> {
636        let mut conn = self.client.clone();
637        let list_key = self.list_key(key);
638
639        let len: i64 = conn
640            .llen(&list_key)
641            .await
642            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
643
644        Ok(len.max(0) as usize)
645    }
646
647    async fn scan(&self, pattern: &str, cursor: u64, limit: usize) -> StorageResult<ScanPage> {
648        let mut conn = self.client.clone();
649        let full_pattern = self.full_key(pattern);
650
651        let (next, raw_keys): (u64, Vec<String>) = redis::cmd("SCAN")
652            .arg(cursor)
653            .arg("MATCH")
654            .arg(&full_pattern)
655            .arg("COUNT")
656            .arg(limit.max(1))
657            .query_async(&mut conn)
658            .await
659            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
660
661        let list_prefix = format!("{}list:", self.key_prefix);
662        let keys: Vec<String> = raw_keys
663            .into_iter()
664            .filter(|k| !k.starts_with(&list_prefix))
665            .filter_map(|k| self.strip_prefix(&k).map(str::to_string))
666            .collect();
667
668        Ok(ScanPage {
669            keys,
670            next_cursor: next,
671        })
672    }
673}
674
675#[cfg(test)]
676mod tests {
677    use super::*;
678    use sa_token_adapter::storage::SaStorage;
679    use std::sync::Arc;
680    use std::time::Duration;
681
682    #[test]
683    fn test_logical_key_stripping_matches_manager_expectations() {
684        let prefix = "phys:";
685        let raw = vec!["phys:sa:token:a".to_string(), "phys:sa:token:b".to_string()];
686        let n = prefix.len();
687        let logical: Vec<String> = raw
688            .into_iter()
689            .map(|k| k.get(n..).map(str::to_string).unwrap_or(k))
690            .collect();
691        assert_eq!(logical, vec!["sa:token:a", "sa:token:b"]);
692    }
693
694    /// 真实 Redis:无 REDIS_URL 时 ignore;测 set/get/ttl/get_del。
695    #[tokio::test]
696    #[ignore = "requires REDIS_URL"]
697    async fn test_redis_set_get_ttl_get_del() {
698        let url = std::env::var("REDIS_URL").expect("REDIS_URL");
699        let storage = RedisStorage::connect(&url).await.expect("connect");
700        let storage: Arc<dyn SaStorage> = Arc::new(storage);
701        let key = format!(
702            "sa:test:gray:{}_{}",
703            std::time::SystemTime::now()
704                .duration_since(std::time::UNIX_EPOCH)
705                .unwrap()
706                .as_nanos(),
707            std::process::id()
708        );
709        storage
710            .set(&key, "v1", Some(Duration::from_secs(30)))
711            .await
712            .expect("set");
713        assert_eq!(storage.get(&key).await.expect("get").as_deref(), Some("v1"));
714        let ttl = storage.ttl(&key).await.expect("ttl");
715        assert!(ttl.is_some());
716        let secs = ttl.unwrap().as_secs();
717        assert!(secs > 0 && secs <= 30);
718        let taken = storage.get_del(&key).await.expect("get_del");
719        assert_eq!(taken.as_deref(), Some("v1"));
720        assert!(storage.get(&key).await.expect("gone").is_none());
721    }
722}