1use 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#[derive(Debug, Clone, Serialize, Deserialize)]
48pub struct RedisConfig {
49 #[serde(default = "default_host")]
51 pub host: String,
52
53 #[serde(default = "default_port")]
55 pub port: u16,
56
57 #[serde(default)]
59 pub password: Option<String>,
60
61 #[serde(default)]
63 pub database: u8,
64
65 #[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 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#[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 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 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 pub fn builder() -> RedisStorageBuilder {
210 RedisStorageBuilder::default()
211 }
212
213 fn full_key(&self, key: &str) -> String {
215 format!("{}{}", self.key_prefix, key)
216 }
217
218 fn list_key(&self, key: &str) -> String {
220 format!("{}list:{}", self.key_prefix, key)
221 }
222
223 fn strip_prefix<'a>(&self, raw: &'a str) -> Option<&'a str> {
225 raw.strip_prefix(&self.key_prefix)
226 }
227
228 pub async fn connect(redis_url: &str) -> StorageResult<Self> {
230 Self::new(redis_url, "").await
231 }
232}
233
234#[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 pub fn host(mut self, host: impl Into<String>) -> Self {
252 self.config.host = host.into();
253 self
254 }
255
256 pub fn port(mut self, port: u16) -> Self {
258 self.config.port = port;
259 self
260 }
261
262 pub fn password(mut self, password: impl Into<String>) -> Self {
264 self.config.password = Some(password.into());
265 self
266 }
267
268 pub fn database(mut self, database: u8) -> Self {
270 self.config.database = database;
271 self
272 }
273
274 pub fn pool_size(mut self, size: u32) -> Self {
276 self.config.pool_size = size;
277 self
278 }
279
280 pub fn key_prefix(mut self, prefix: impl Into<String>) -> Self {
282 self.key_prefix = Some(prefix.into());
283 self
284 }
285
286 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), -1 => Ok(None), 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 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 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 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 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 #[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}