1use 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
14const SHARD_COUNT: usize = 16;
17const SHARD_MASK: usize = SHARD_COUNT - 1;
18
19#[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#[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#[derive(Debug, Default)]
61struct MemoryState {
62 scalars: HashMap<String, StorageItem>,
63 lists: HashMap<String, ListItem>,
64}
65
66fn 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#[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#[derive(Debug, Clone)]
102pub struct MemoryStorage {
103 shards: Arc<[RwLock<MemoryState>]>,
104}
105
106impl MemoryStorage {
107 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 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 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 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 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}