1use super::StoreError;
4use std::collections::HashMap;
5use std::sync::Mutex;
6
7#[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 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#[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 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 fn revoke_subject(&self, subject: &str) -> Result<usize, StoreError>;
40 fn recent_logins(&self, subject: &str) -> Result<Vec<LoginPoint>, StoreError>;
42 fn record_login(&self, subject: &str, point: LoginPoint) -> Result<(), StoreError>;
43 fn purge_expired(&self, now: u64) -> Result<usize, StoreError>;
45}
46
47pub const MAX_LOGINS_PER_SUBJECT: usize = 10;
49
50#[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 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 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 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 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 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 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 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 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}