Skip to main content

sa_token_storage_memory/
lib.rs

1// Author: 金书记
2//
3//! 内存存储实现(开发/单机/无持久化场景)
4
5use std::collections::HashMap;
6use std::sync::Arc;
7use std::time::Duration;
8
9use async_trait::async_trait;
10use chrono::{DateTime, Utc};
11use sa_token_adapter::storage::{SaStorage, ScanPage, StorageError, StorageResult};
12use tokio::sync::RwLock;
13
14/// 分片数:2 的幂,用位与代替取模。同一 key 的标量与 list 必须落在同一分片。
15/// Shard count (power of two). Scalar and list for the same key must share a shard.
16const SHARD_COUNT: usize = 16;
17const SHARD_MASK: usize = SHARD_COUNT - 1;
18
19/// 标量键值项
20#[derive(Debug, Clone)]
21struct StorageItem {
22    value: String,
23    expire_at: Option<DateTime<Utc>>,
24}
25
26impl StorageItem {
27    fn new(value: String, ttl: Option<Duration>) -> Self {
28        let expire_at = ttl
29            .and_then(|d| chrono::Duration::from_std(d).ok())
30            .map(|d| Utc::now() + d);
31        Self { value, expire_at }
32    }
33
34    fn is_expired(&self) -> bool {
35        self.expire_at.is_some_and(|t| Utc::now() > t)
36    }
37}
38
39/// 集合键项
40#[derive(Debug, Clone, Default)]
41struct ListItem {
42    members: Vec<String>,
43    expire_at: Option<DateTime<Utc>>,
44}
45
46impl ListItem {
47    fn touch_ttl(&mut self, ttl: Option<Duration>) {
48        if let Some(d) = ttl.and_then(|d| chrono::Duration::from_std(d).ok()) {
49            self.expire_at = Some(Utc::now() + d);
50        }
51    }
52
53    fn is_expired(&self) -> bool {
54        self.expire_at.is_some_and(|t| Utc::now() > t)
55    }
56}
57
58/// 单分片状态(同 key 的 scalar / list 同锁)
59/// Per-shard state (scalar and list for one key share the lock)
60#[derive(Debug, Default)]
61struct MemoryState {
62    scalars: HashMap<String, StorageItem>,
63    lists: HashMap<String, ListItem>,
64}
65
66/// 将 glob 风格 pattern 转为锚定正则(转义元字符,`*` → `.*`)
67///
68/// 修复 1.2-E:旧版 `pattern.replace("*", ".*")` 未锚定,导致 `sa:token:*` 误匹配 `x-sa:token:y`
69fn glob_to_regex(pattern: &str) -> Result<regex::Regex, StorageError> {
70    let mut re = String::from("^");
71    for ch in pattern.chars() {
72        match ch {
73            '*' => re.push_str(".*"),
74            '?' => re.push('.'),
75            c if ".^$+{}[]|()\\".contains(c) => {
76                re.push('\\');
77                re.push(c);
78            }
79            c => re.push(c),
80        }
81    }
82    re.push('$');
83    regex::Regex::new(&re)
84        .map_err(|e| StorageError::OperationFailed(format!("Invalid pattern: {e}")))
85}
86
87/// 内部可控 key 用 FNV-1a;不引入 ahash。
88/// FNV-1a for library-controlled keys; no ahash dependency.
89#[inline]
90fn shard_index(key: &str) -> usize {
91    let mut hash = 0xcbf29ce484222325u64;
92    for byte in key.as_bytes() {
93        hash ^= u64::from(*byte);
94        hash = hash.wrapping_mul(0x100000001b3);
95    }
96    (hash as usize) & SHARD_MASK
97}
98
99/// In-memory [`SaStorage`] for tests and single-node use.
100/// 进程内内存存储,适用于测试与单机场景。
101#[derive(Debug, Clone)]
102pub struct MemoryStorage {
103    shards: Arc<[RwLock<MemoryState>]>,
104}
105
106impl MemoryStorage {
107    /// Create an empty memory store | 创建空的内存存储
108    pub fn new() -> Self {
109        let shards: Vec<_> = (0..SHARD_COUNT)
110            .map(|_| RwLock::new(MemoryState::default()))
111            .collect();
112        Self {
113            shards: shards.into(),
114        }
115    }
116
117    #[inline]
118    fn shard(&self, key: &str) -> &RwLock<MemoryState> {
119        // `shard_index` is masked to `0..SHARD_COUNT`; get avoids clippy::indexing_slicing.
120        // `shard_index` 经掩码落在 `0..SHARD_COUNT`;用 get 规避 indexing_slicing。
121        match self.shards.get(shard_index(key)) {
122            Some(s) => s,
123            None => unreachable!("shard_index always in 0..SHARD_COUNT"),
124        }
125    }
126
127    /// 清理过期标量与集合(逐分片写锁,禁止同时持多把写锁)
128    /// Drop expired scalars/lists (one write lock at a time; never hold multiple)
129    pub async fn cleanup_expired(&self) {
130        for shard in self.shards.iter() {
131            let mut s = shard.write().await;
132            s.scalars.retain(|_, v| !v.is_expired());
133            s.lists.retain(|_, v| !v.is_expired());
134        }
135    }
136}
137
138impl Default for MemoryStorage {
139    fn default() -> Self {
140        Self::new()
141    }
142}
143
144#[async_trait]
145impl SaStorage for MemoryStorage {
146    async fn get(&self, key: &str) -> StorageResult<Option<String>> {
147        let s = self.shard(key).read().await;
148        match s.scalars.get(key) {
149            Some(item) if !item.is_expired() => Ok(Some(item.value.clone())),
150            Some(_) => {
151                drop(s);
152                self.delete(key).await?;
153                Ok(None)
154            }
155            None => Ok(None),
156        }
157    }
158
159    async fn set(&self, key: &str, value: &str, ttl: Option<Duration>) -> StorageResult<()> {
160        let mut s = self.shard(key).write().await;
161        s.scalars
162            .insert(key.to_string(), StorageItem::new(value.to_string(), ttl));
163        Ok(())
164    }
165
166    async fn delete(&self, key: &str) -> StorageResult<()> {
167        let mut s = self.shard(key).write().await;
168        s.scalars.remove(key);
169        s.lists.remove(key);
170        Ok(())
171    }
172
173    async fn exists(&self, key: &str) -> StorageResult<bool> {
174        let s = self.shard(key).read().await;
175        if let Some(item) = s.scalars.get(key) {
176            return Ok(!item.is_expired());
177        }
178        if let Some(list) = s.lists.get(key) {
179            return Ok(!list.is_expired() && !list.members.is_empty());
180        }
181        Ok(false)
182    }
183
184    async fn expire(&self, key: &str, ttl: Duration) -> StorageResult<()> {
185        let mut s = self.shard(key).write().await;
186        let Some(delta) = chrono::Duration::from_std(ttl).ok() else {
187            return Ok(());
188        };
189        let exp = Utc::now() + delta;
190        if let Some(item) = s.scalars.get_mut(key) {
191            item.expire_at = Some(exp);
192        }
193        if let Some(list) = s.lists.get_mut(key) {
194            list.expire_at = Some(exp);
195        }
196        Ok(())
197    }
198
199    async fn ttl(&self, key: &str) -> StorageResult<Option<Duration>> {
200        let s = self.shard(key).read().await;
201        let expire_at = s
202            .scalars
203            .get(key)
204            .and_then(|i| i.expire_at)
205            .or_else(|| s.lists.get(key).and_then(|l| l.expire_at));
206        match expire_at {
207            Some(exp) if exp > Utc::now() => Ok(Some(
208                (exp - Utc::now())
209                    .to_std()
210                    .map_err(|e| StorageError::InternalError(e.to_string()))?,
211            )),
212            Some(_) => Ok(Some(Duration::ZERO)),
213            None => Ok(None),
214        }
215    }
216
217    /// 批量读:按 key 分片取,不跨分片持锁。
218    /// Batched get: lock per key's shard; never hold multiple shard locks.
219    async fn mget(&self, keys: &[&str]) -> StorageResult<Vec<Option<String>>> {
220        let mut results = Vec::with_capacity(keys.len());
221        for &key in keys {
222            results.push(self.get(key).await?);
223        }
224        Ok(results)
225    }
226
227    async fn mset(&self, items: &[(&str, &str)], ttl: Option<Duration>) -> StorageResult<()> {
228        for (key, value) in items {
229            self.set(key, value, ttl).await?;
230        }
231        Ok(())
232    }
233
234    async fn mdel(&self, keys: &[&str]) -> StorageResult<()> {
235        for key in keys {
236            self.delete(key).await?;
237        }
238        Ok(())
239    }
240
241    async fn incr(&self, key: &str) -> StorageResult<i64> {
242        let mut s = self.shard(key).write().await;
243        let current = s
244            .scalars
245            .get(key)
246            .filter(|item| !item.is_expired())
247            .and_then(|item| item.value.parse::<i64>().ok())
248            .unwrap_or(0);
249        let new_value = current + 1;
250        s.scalars.insert(
251            key.to_string(),
252            StorageItem::new(new_value.to_string(), None),
253        );
254        Ok(new_value)
255    }
256
257    async fn decr(&self, key: &str) -> StorageResult<i64> {
258        let mut s = self.shard(key).write().await;
259        let current = s
260            .scalars
261            .get(key)
262            .filter(|item| !item.is_expired())
263            .and_then(|item| item.value.parse::<i64>().ok())
264            .unwrap_or(0);
265        let new_value = current - 1;
266        s.scalars.insert(
267            key.to_string(),
268            StorageItem::new(new_value.to_string(), None),
269        );
270        Ok(new_value)
271    }
272
273    async fn clear(&self) -> StorageResult<()> {
274        for shard in self.shards.iter() {
275            let mut s = shard.write().await;
276            s.scalars.clear();
277            s.lists.clear();
278        }
279        Ok(())
280    }
281
282    async fn set_if_absent(
283        &self,
284        key: &str,
285        value: &str,
286        ttl: Option<Duration>,
287    ) -> StorageResult<bool> {
288        let mut s = self.shard(key).write().await;
289        if let Some(existing) = s.scalars.get(key) {
290            if !existing.is_expired() {
291                return Ok(false);
292            }
293            s.scalars.remove(key);
294        }
295        s.scalars
296            .insert(key.to_string(), StorageItem::new(value.to_string(), ttl));
297        Ok(true)
298    }
299
300    async fn get_del(&self, key: &str) -> StorageResult<Option<String>> {
301        let mut s = self.shard(key).write().await;
302        match s.scalars.remove(key) {
303            Some(item) if !item.is_expired() => Ok(Some(item.value)),
304            _ => Ok(None),
305        }
306    }
307
308    async fn compare_and_swap(
309        &self,
310        key: &str,
311        expected: Option<&str>,
312        new_value: &str,
313        ttl: Option<Duration>,
314    ) -> StorageResult<bool> {
315        let mut s = self.shard(key).write().await;
316        let current = s
317            .scalars
318            .get(key)
319            .filter(|i| !i.is_expired())
320            .map(|i| i.value.as_str());
321        if current != expected {
322            return Ok(false);
323        }
324        s.scalars.insert(
325            key.to_string(),
326            StorageItem::new(new_value.to_string(), ttl),
327        );
328        Ok(true)
329    }
330
331    async fn compare_and_delete(&self, key: &str, expected: &str) -> StorageResult<bool> {
332        let mut s = self.shard(key).write().await;
333        let ok = s
334            .scalars
335            .get(key)
336            .filter(|i| !i.is_expired())
337            .is_some_and(|i| i.value == expected);
338        if ok {
339            s.scalars.remove(key);
340        }
341        Ok(ok)
342    }
343
344    async fn list_push(
345        &self,
346        key: &str,
347        member: &str,
348        unique: bool,
349        ttl: Option<Duration>,
350    ) -> StorageResult<usize> {
351        let mut s = self.shard(key).write().await;
352        let entry = s.lists.entry(key.to_string()).or_default();
353        if entry.is_expired() {
354            entry.members.clear();
355            entry.expire_at = None;
356        }
357        if unique && entry.members.iter().any(|m| m == member) {
358            entry.touch_ttl(ttl);
359            return Ok(entry.members.len());
360        }
361        entry.members.push(member.to_string());
362        entry.touch_ttl(ttl);
363        Ok(entry.members.len())
364    }
365
366    async fn list_remove(&self, key: &str, member: &str) -> StorageResult<bool> {
367        let mut s = self.shard(key).write().await;
368        let Some(entry) = s.lists.get_mut(key) else {
369            return Ok(false);
370        };
371        if entry.is_expired() {
372            entry.members.clear();
373            return Ok(false);
374        }
375        let before = entry.members.len();
376        entry.members.retain(|m| m != member);
377        Ok(entry.members.len() < before)
378    }
379
380    async fn list_range(
381        &self,
382        key: &str,
383        start: usize,
384        limit: Option<usize>,
385    ) -> StorageResult<Vec<String>> {
386        let s = self.shard(key).read().await;
387        let Some(entry) = s.lists.get(key) else {
388            return Ok(Vec::new());
389        };
390        if entry.is_expired() {
391            return Ok(Vec::new());
392        }
393        let slice: Vec<String> = entry
394            .members
395            .iter()
396            .skip(start)
397            .take(limit.unwrap_or(usize::MAX))
398            .cloned()
399            .collect();
400        Ok(slice)
401    }
402
403    async fn list_len(&self, key: &str) -> StorageResult<usize> {
404        let s = self.shard(key).read().await;
405        Ok(s.lists
406            .get(key)
407            .filter(|l| !l.is_expired())
408            .map(|l| l.members.len())
409            .unwrap_or(0))
410    }
411
412    async fn scan(&self, pattern: &str, cursor: u64, limit: usize) -> StorageResult<ScanPage> {
413        let re = glob_to_regex(pattern)?;
414        // 逐分片读锁收集,禁止同时持多把锁。
415        // Collect under one shard read lock at a time.
416        let mut keys: Vec<String> = Vec::new();
417        for shard in self.shards.iter() {
418            let s = shard.read().await;
419            for (k, v) in s.scalars.iter() {
420                if !v.is_expired() {
421                    keys.push(k.clone());
422                }
423            }
424        }
425        keys.sort();
426        keys.dedup();
427        keys.retain(|k| re.is_match(k));
428        let start = cursor as usize;
429        if start >= keys.len() {
430            return Ok(ScanPage {
431                keys: Vec::new(),
432                next_cursor: 0,
433            });
434        }
435        let end = (start + limit).min(keys.len());
436        let page_keys = keys.get(start..end).unwrap_or(&[]).to_vec();
437        let next = if end >= keys.len() { 0 } else { end as u64 };
438        Ok(ScanPage {
439            keys: page_keys,
440            next_cursor: next,
441        })
442    }
443}
444
445#[cfg(test)]
446mod tests {
447    use super::*;
448    use sa_token_adapter::CountingStorage;
449    use std::sync::Arc;
450
451    #[tokio::test]
452    async fn test_set_if_absent_and_get_del() {
453        let storage = MemoryStorage::new();
454        assert!(storage.set_if_absent("n1", "v", None).await.unwrap());
455        assert!(!storage.set_if_absent("n1", "v2", None).await.unwrap());
456        assert_eq!(storage.get_del("n1").await.unwrap(), Some("v".into()));
457        assert_eq!(storage.get_del("n1").await.unwrap(), None);
458    }
459
460    #[tokio::test]
461    async fn test_concurrent_nonce_get_del() {
462        let storage = Arc::new(MemoryStorage::new());
463        storage.set_if_absent("nonce:x", "1", None).await.unwrap();
464        let mut got = 0usize;
465        for _ in 0..100 {
466            if storage.get_del("nonce:x").await.unwrap().is_some() {
467                got += 1;
468            }
469        }
470        assert_eq!(got, 1);
471    }
472
473    #[tokio::test]
474    async fn test_compare_and_swap_and_delete() {
475        let storage = MemoryStorage::new();
476        storage.set("k", "old", None).await.unwrap();
477        assert!(
478            !storage
479                .compare_and_swap("k", Some("wrong"), "new", None)
480                .await
481                .unwrap()
482        );
483        assert_eq!(storage.get("k").await.unwrap(), Some("old".into()));
484        assert!(
485            storage
486                .compare_and_swap("k", Some("old"), "new", None)
487                .await
488                .unwrap()
489        );
490        assert_eq!(storage.get("k").await.unwrap(), Some("new".into()));
491        assert!(!storage.compare_and_delete("k", "wrong").await.unwrap());
492        assert!(storage.compare_and_delete("k", "new").await.unwrap());
493        assert!(!storage.exists("k").await.unwrap());
494    }
495
496    #[tokio::test]
497    async fn test_concurrent_list_push() {
498        let storage = Arc::new(MemoryStorage::new());
499        let mut handles = Vec::new();
500        for i in 0..50 {
501            let s = Arc::clone(&storage);
502            handles.push(tokio::spawn(async move {
503                s.list_push("idx", &format!("t{i}"), false, None)
504                    .await
505                    .unwrap();
506            }));
507        }
508        for h in handles {
509            h.await.unwrap();
510        }
511        assert_eq!(storage.list_len("idx").await.unwrap(), 50);
512    }
513
514    #[tokio::test]
515    async fn test_scan_pagination_complete() {
516        let storage = MemoryStorage::new();
517        for i in 0..1000 {
518            storage
519                .set(&format!("sa:token:{i:04}"), "v", None)
520                .await
521                .unwrap();
522        }
523        let mut cursor = 0u64;
524        let mut all = Vec::new();
525        loop {
526            let page = storage.scan("sa:token:*", cursor, 100).await.unwrap();
527            all.extend(page.keys);
528            if page.next_cursor == 0 {
529                break;
530            }
531            cursor = page.next_cursor;
532        }
533        all.sort();
534        all.dedup();
535        assert_eq!(all.len(), 1000);
536    }
537
538    #[tokio::test]
539    async fn test_list_push_len_and_scan_anchor() {
540        let storage = MemoryStorage::new();
541        for i in 0..50 {
542            storage
543                .list_push("idx", &format!("t{i}"), false, None)
544                .await
545                .unwrap();
546        }
547        assert_eq!(storage.list_len("idx").await.unwrap(), 50);
548        storage.set("x-sa:token:y", "bad", None).await.unwrap();
549        storage.set("sa:token:a", "ok", None).await.unwrap();
550        let page = storage.scan("sa:token:*", 0, 100).await.unwrap();
551        assert_eq!(page.keys, vec!["sa:token:a".to_string()]);
552    }
553
554    #[tokio::test]
555    async fn test_counting_storage_decorator() {
556        let inner = Arc::new(MemoryStorage::new()) as Arc<dyn SaStorage>;
557        let counting = CountingStorage::new(Arc::clone(&inner));
558
559        counting.set("k", "v", None).await.unwrap();
560        assert_eq!(counting.get("k").await.unwrap(), Some("v".into()));
561        assert_eq!(counting.get_count(), 1);
562        assert_eq!(counting.set_count(), 1);
563
564        counting.reset_counts();
565        counting.get("k").await.unwrap();
566        counting.delete("k").await.unwrap();
567        assert_eq!(counting.get_count(), 1);
568        assert_eq!(counting.delete_count(), 1);
569        assert_eq!(counting.set_count(), 0);
570    }
571}