Skip to main content

security_rust/session/
store.rs

1// Copyright (c) 2026 erik <erik@erik.xyz> — https://erik.xyz
2
3use super::StoreError;
4use std::collections::HashMap;
5use std::sync::Mutex;
6
7/// 会话记录 —— 这张表就是「token 会话」。
8#[derive(Debug, Clone, PartialEq)]
9pub struct SessionRecord {
10    pub token: String,
11    pub subject: String,
12    pub fingerprint: String,
13    pub location: Option<String>,
14    pub coords: Option<(f64, f64)>,
15    /// 登录时的签名基线。
16    pub signature: Option<String>,
17    pub issued_at: u64,
18    pub last_seen: u64,
19    pub expires_at: u64,
20    pub revoked: bool,
21}
22
23/// 一次登录的位置快照,用于异地检测与不可能旅行。
24#[derive(Debug, Clone, PartialEq)]
25pub struct LoginPoint {
26    pub location: Option<String>,
27    pub coords: Option<(f64, f64)>,
28    pub at: u64,
29}
30
31pub trait SessionStore: Send + Sync {
32    fn put(&self, rec: SessionRecord) -> Result<(), StoreError>;
33    /// 原样返回记录,**不过滤过期**。过期判定归 `guard` ——
34    /// 若此处对过期记录返回 `None`,`verify` 将无法区分 `TokenExpired` 与 `TokenUnknown`。
35    fn get(&self, token: &str) -> Result<Option<SessionRecord>, StoreError>;
36    fn touch(&self, token: &str, now: u64) -> Result<(), StoreError>;
37    fn revoke(&self, token: &str) -> Result<(), StoreError>;
38    /// 吊销某 subject 的全部会话,返回受影响条数。
39    fn revoke_subject(&self, subject: &str) -> Result<usize, StoreError>;
40    /// 该 subject 的登录历史,按时间升序(最旧在前)。
41    fn recent_logins(&self, subject: &str) -> Result<Vec<LoginPoint>, StoreError>;
42    fn record_login(&self, subject: &str, point: LoginPoint) -> Result<(), StoreError>;
43    /// 清除已过期记录,返回清除条数。
44    fn purge_expired(&self, now: u64) -> Result<usize, StoreError>;
45}
46
47/// 每个 subject 保留的登录历史条数上限。
48pub const MAX_LOGINS_PER_SUBJECT: usize = 10;
49
50/// 内存后端。无后台线程 —— 过期判定归 guard,内存回收靠 `purge_expired`。
51#[derive(Debug)]
52pub struct MemoryStore {
53    sessions: Mutex<HashMap<String, SessionRecord>>,
54    logins: Mutex<HashMap<String, Vec<LoginPoint>>>,
55}
56
57impl Default for MemoryStore {
58    fn default() -> Self {
59        Self::new()
60    }
61}
62
63impl MemoryStore {
64    pub fn new() -> Self {
65        Self {
66            sessions: Mutex::new(HashMap::new()),
67            logins: Mutex::new(HashMap::new()),
68        }
69    }
70
71    /// 互斥锁获取。线程 panic 导致锁中毒时恢复内部数据而非永久 Err:
72    /// 本 store 的每个操作都是单次 HashMap 读/写,不存在「改到一半」的不变量,
73    /// 恢复是安全的;而让一次 panic 永久锁死整个会话存储是一种自我 DoS。
74    fn lock<T>(m: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
75        m.lock().unwrap_or_else(|e| e.into_inner())
76    }
77}
78
79impl SessionStore for MemoryStore {
80    fn put(&self, rec: SessionRecord) -> Result<(), StoreError> {
81        Self::lock(&self.sessions).insert(rec.token.clone(), rec);
82        Ok(())
83    }
84
85    fn get(&self, token: &str) -> Result<Option<SessionRecord>, StoreError> {
86        Ok(Self::lock(&self.sessions).get(token).cloned())
87    }
88
89    fn touch(&self, token: &str, now: u64) -> Result<(), StoreError> {
90        if let Some(r) = Self::lock(&self.sessions).get_mut(token) {
91            r.last_seen = now;
92        }
93        Ok(())
94    }
95
96    fn revoke(&self, token: &str) -> Result<(), StoreError> {
97        if let Some(r) = Self::lock(&self.sessions).get_mut(token) {
98            r.revoked = true;
99        }
100        Ok(())
101    }
102
103    // ponytail: 持全局锁的 O(n) 全表扫描。内存后端规模下可接受;
104    // 若单 subject 会话数上到万级或需跨实例,加 subject -> tokens 索引。
105    fn revoke_subject(&self, subject: &str) -> Result<usize, StoreError> {
106        let mut n = 0;
107        for r in Self::lock(&self.sessions).values_mut() {
108            if r.subject == subject && !r.revoked {
109                r.revoked = true;
110                n += 1;
111            }
112        }
113        Ok(n)
114    }
115
116    fn recent_logins(&self, subject: &str) -> Result<Vec<LoginPoint>, StoreError> {
117        Ok(Self::lock(&self.logins)
118            .get(subject)
119            .cloned()
120            .unwrap_or_default())
121    }
122
123    fn record_login(&self, subject: &str, point: LoginPoint) -> Result<(), StoreError> {
124        let mut g = Self::lock(&self.logins);
125        let v = g.entry(subject.to_string()).or_default();
126        v.push(point);
127        // 有界:只保留最近 MAX_LOGINS_PER_SUBJECT 条
128        if v.len() > MAX_LOGINS_PER_SUBJECT {
129            v.drain(..v.len() - MAX_LOGINS_PER_SUBJECT);
130        }
131        Ok(())
132    }
133
134    fn purge_expired(&self, now: u64) -> Result<usize, StoreError> {
135        let mut g = Self::lock(&self.sessions);
136        let before = g.len();
137        g.retain(|_, r| r.expires_at > now);
138        Ok(before - g.len())
139    }
140}
141
142#[cfg(test)]
143mod tests {
144    use super::*;
145
146    fn rec(token: &str, subject: &str, expires_at: u64) -> SessionRecord {
147        SessionRecord {
148            token: token.into(),
149            subject: subject.into(),
150            fingerprint: "ip=1.2.3.4|ua=curl".into(),
151            location: Some("CN-BJ".into()),
152            coords: Some((39.9042, 116.4074)),
153            signature: None,
154            issued_at: 1_000,
155            last_seen: 1_000,
156            expires_at,
157            revoked: false,
158        }
159    }
160
161    #[test]
162    fn put_then_get_roundtrip() {
163        let s = MemoryStore::new();
164        s.put(rec("t1", "u1", 2_000)).unwrap();
165        let got = s.get("t1").unwrap().expect("present");
166        assert_eq!(got.token, "t1");
167        assert_eq!(got.subject, "u1");
168    }
169
170    #[test]
171    fn get_unknown_token_is_none() {
172        let s = MemoryStore::new();
173        assert!(s.get("nope").unwrap().is_none());
174    }
175
176    #[test]
177    fn get_returns_expired_record_unfiltered() {
178        // 关键契约:store 不做过期过滤,过期判定归 guard。
179        // 否则 verify 无法区分 TokenExpired 与 TokenUnknown。
180        let s = MemoryStore::new();
181        s.put(rec("t1", "u1", 500)).unwrap();
182        let got = s.get("t1").unwrap().expect("must still be returned");
183        assert_eq!(got.expires_at, 500);
184    }
185
186    #[test]
187    fn touch_updates_last_seen_only() {
188        let s = MemoryStore::new();
189        s.put(rec("t1", "u1", 9_999)).unwrap();
190        s.touch("t1", 1_500).unwrap();
191        let got = s.get("t1").unwrap().unwrap();
192        assert_eq!(got.last_seen, 1_500);
193        assert_eq!(got.issued_at, 1_000);
194        assert_eq!(got.expires_at, 9_999);
195    }
196
197    #[test]
198    fn touch_unknown_token_is_ok() {
199        let s = MemoryStore::new();
200        assert!(s.touch("nope", 1_500).is_ok());
201    }
202
203    #[test]
204    fn revoke_marks_record_not_deletes_it() {
205        let s = MemoryStore::new();
206        s.put(rec("t1", "u1", 9_999)).unwrap();
207        s.revoke("t1").unwrap();
208        let got = s.get("t1").unwrap().expect("kept for diagnosis");
209        assert!(got.revoked);
210    }
211
212    #[test]
213    fn revoke_unknown_token_is_ok() {
214        let s = MemoryStore::new();
215        assert!(s.revoke("nope").is_ok());
216    }
217
218    #[test]
219    fn revoke_subject_counts_only_newly_revoked() {
220        let s = MemoryStore::new();
221        s.put(rec("t1", "u1", 9_999)).unwrap();
222        s.put(rec("t2", "u1", 9_999)).unwrap();
223        s.put(rec("t3", "u2", 9_999)).unwrap();
224        assert_eq!(s.revoke_subject("u1").unwrap(), 2);
225        // 再调一次不该重复计数
226        assert_eq!(s.revoke_subject("u1").unwrap(), 0);
227        assert!(!s.get("t3").unwrap().unwrap().revoked);
228    }
229
230    #[test]
231    fn revoke_subject_unknown_is_zero() {
232        let s = MemoryStore::new();
233        assert_eq!(s.revoke_subject("ghost").unwrap(), 0);
234    }
235
236    fn point(loc: &str, at: u64) -> LoginPoint {
237        LoginPoint {
238            location: Some(loc.into()),
239            coords: None,
240            at,
241        }
242    }
243
244    #[test]
245    fn login_history_is_ascending() {
246        let s = MemoryStore::new();
247        s.record_login("u1", point("CN-BJ", 100)).unwrap();
248        s.record_login("u1", point("CN-SH", 200)).unwrap();
249        let h = s.recent_logins("u1").unwrap();
250        assert_eq!(h.len(), 2);
251        assert_eq!(h[0].at, 100);
252        assert_eq!(h[1].at, 200);
253    }
254
255    #[test]
256    fn login_history_is_bounded() {
257        let s = MemoryStore::new();
258        for i in 0..(MAX_LOGINS_PER_SUBJECT as u64 + 5) {
259            s.record_login("u1", point("CN-BJ", 100 + i)).unwrap();
260        }
261        let h = s.recent_logins("u1").unwrap();
262        assert_eq!(h.len(), MAX_LOGINS_PER_SUBJECT);
263        // 保留的是最近的:最旧的一条应为 at = 105
264        assert_eq!(h[0].at, 105);
265    }
266
267    #[test]
268    fn recent_logins_unknown_subject_is_empty() {
269        let s = MemoryStore::new();
270        assert!(s.recent_logins("ghost").unwrap().is_empty());
271    }
272
273    #[test]
274    fn purge_expired_removes_only_expired() {
275        let s = MemoryStore::new();
276        s.put(rec("old", "u1", 500)).unwrap();
277        s.put(rec("new", "u1", 5_000)).unwrap();
278        assert_eq!(s.purge_expired(1_000).unwrap(), 1);
279        assert!(s.get("old").unwrap().is_none());
280        assert!(s.get("new").unwrap().is_some());
281    }
282
283    #[test]
284    fn purge_expired_boundary_is_exclusive() {
285        // expires_at == now 视为已过期
286        let s = MemoryStore::new();
287        s.put(rec("t", "u1", 1_000)).unwrap();
288        assert_eq!(s.purge_expired(1_000).unwrap(), 1);
289    }
290
291    #[test]
292    fn purge_expired_on_empty_is_zero() {
293        let s = MemoryStore::new();
294        assert_eq!(s.purge_expired(1_000).unwrap(), 0);
295    }
296
297    #[test]
298    fn lock_recovers_from_poisoned_mutex() {
299        // 钉住 lock() 的恢复不变量:一次 panic 不能永久锁死会话存储。
300        let m = Mutex::new(rec("t1", "u1", 9_999));
301        std::panic::catch_unwind(|| {
302            let _guard = m.lock().unwrap();
303            panic!("poison");
304        })
305        .unwrap_err();
306        assert!(m.is_poisoned());
307
308        let g = MemoryStore::lock(&m);
309        assert_eq!(g.token, "t1");
310        assert_eq!(g.expires_at, 9_999);
311    }
312}