shell_tunnel/session/
store.rs1use std::collections::HashMap;
4use std::sync::RwLock;
5use std::time::Instant;
6
7use super::{SessionContext, SessionId, SessionState};
8use crate::error::ShellTunnelError;
9use crate::Result;
10
11#[derive(Debug, Clone)]
19pub struct Session {
20 pub id: SessionId,
22 pub state: SessionState,
24 pub context: SessionContext,
26 pub created_at: Instant,
28 pub last_activity: Instant,
30}
31
32impl Session {
33 pub fn new(id: SessionId) -> Self {
35 let now = Instant::now();
36
37 Self {
38 id,
39 state: SessionState::Created,
40 context: SessionContext::new(),
41 created_at: now,
42 last_activity: now,
43 }
44 }
45
46 pub fn touch(&mut self) {
48 self.last_activity = Instant::now();
49 }
50
51 pub fn idle_duration(&self) -> std::time::Duration {
53 self.last_activity.elapsed()
54 }
55}
56
57pub struct SessionStore {
59 sessions: RwLock<HashMap<SessionId, Session>>,
60}
61
62impl SessionStore {
63 pub fn new() -> Self {
65 Self {
66 sessions: RwLock::new(HashMap::new()),
67 }
68 }
69
70 pub fn create(&self) -> Result<SessionId> {
74 let id = SessionId::new();
75 let session = Session::new(id);
76
77 let mut sessions = self
78 .sessions
79 .write()
80 .map_err(|_| ShellTunnelError::LockPoisoned)?;
81
82 sessions.insert(id, session);
83 Ok(id)
84 }
85
86 pub fn get(&self, id: &SessionId) -> Result<Option<Session>> {
88 let sessions = self
89 .sessions
90 .read()
91 .map_err(|_| ShellTunnelError::LockPoisoned)?;
92 Ok(sessions.get(id).cloned())
93 }
94
95 pub fn contains(&self, id: &SessionId) -> Result<bool> {
97 let sessions = self
98 .sessions
99 .read()
100 .map_err(|_| ShellTunnelError::LockPoisoned)?;
101 Ok(sessions.contains_key(id))
102 }
103
104 pub fn update<F>(&self, id: &SessionId, f: F) -> Result<()>
109 where
110 F: FnOnce(&mut Session),
111 {
112 let mut sessions = self
113 .sessions
114 .write()
115 .map_err(|_| ShellTunnelError::LockPoisoned)?;
116
117 let session = sessions
118 .get_mut(id)
119 .ok_or_else(|| ShellTunnelError::SessionNotFound(id.to_string()))?;
120
121 f(session);
122 Ok(())
123 }
124
125 pub fn remove(&self, id: &SessionId) -> Result<Option<Session>> {
129 let mut sessions = self
130 .sessions
131 .write()
132 .map_err(|_| ShellTunnelError::LockPoisoned)?;
133 Ok(sessions.remove(id))
134 }
135
136 pub fn count(&self) -> usize {
138 self.sessions.read().map(|s| s.len()).unwrap_or(0)
139 }
140
141 pub fn list_ids(&self) -> Result<Vec<SessionId>> {
143 let sessions = self
144 .sessions
145 .read()
146 .map_err(|_| ShellTunnelError::LockPoisoned)?;
147 Ok(sessions.keys().copied().collect())
148 }
149
150 pub fn remove_matching<F>(&self, predicate: F) -> Result<usize>
154 where
155 F: Fn(&Session) -> bool,
156 {
157 let mut sessions = self
158 .sessions
159 .write()
160 .map_err(|_| ShellTunnelError::LockPoisoned)?;
161
162 let before = sessions.len();
163 sessions.retain(|_, session| !predicate(session));
164 Ok(before - sessions.len())
165 }
166
167 pub fn sweep_idle(&self, ttl: std::time::Duration) -> Result<Vec<SessionId>> {
202 let mut sessions = self
203 .sessions
204 .write()
205 .map_err(|_| ShellTunnelError::LockPoisoned)?;
206
207 let expired: Vec<SessionId> = sessions
208 .values()
209 .filter(|s| s.state != SessionState::Active && s.idle_duration() > ttl)
210 .map(|s| s.id)
211 .collect();
212
213 for id in &expired {
214 sessions.remove(id);
215 }
216 Ok(expired)
217 }
218}
219
220impl Default for SessionStore {
221 fn default() -> Self {
222 Self::new()
223 }
224}
225
226#[cfg(test)]
227mod tests {
228 use super::*;
229
230 #[test]
231 fn test_create_session() {
232 let store = SessionStore::new();
233 let id = store.create().unwrap();
234
235 assert!(store.contains(&id).unwrap());
236 assert_eq!(store.count(), 1);
237 }
238
239 #[test]
240 fn test_get_session() {
241 let store = SessionStore::new();
242 let id = store.create().unwrap();
243
244 let session = store.get(&id).unwrap().unwrap();
245 assert_eq!(session.id, id);
246 assert_eq!(session.state, SessionState::Created);
247 }
248
249 #[test]
250 fn test_get_nonexistent() {
251 let store = SessionStore::new();
252 let fake_id = SessionId::from_raw(999999);
253
254 let result = store.get(&fake_id).unwrap();
255 assert!(result.is_none());
256 }
257
258 #[test]
259 fn test_update_session() {
260 let store = SessionStore::new();
261 let id = store.create().unwrap();
262
263 store
264 .update(&id, |s| {
265 s.state = SessionState::Active;
266 })
267 .unwrap();
268
269 let session = store.get(&id).unwrap().unwrap();
270 assert_eq!(session.state, SessionState::Active);
271 }
272
273 #[test]
274 fn test_update_nonexistent() {
275 let store = SessionStore::new();
276 let fake_id = SessionId::from_raw(999999);
277
278 let result = store.update(&fake_id, |_| {});
279 assert!(result.is_err());
280 }
281
282 #[test]
283 fn test_remove_session() {
284 let store = SessionStore::new();
285 let id = store.create().unwrap();
286
287 let removed = store.remove(&id).unwrap();
288 assert!(removed.is_some());
289 assert_eq!(removed.unwrap().id, id);
290
291 assert!(!store.contains(&id).unwrap());
292 assert_eq!(store.count(), 0);
293 }
294
295 #[test]
296 fn test_list_ids() {
297 let store = SessionStore::new();
298 let id1 = store.create().unwrap();
299 let id2 = store.create().unwrap();
300 let id3 = store.create().unwrap();
301
302 let ids = store.list_ids().unwrap();
303 assert_eq!(ids.len(), 3);
304 assert!(ids.contains(&id1));
305 assert!(ids.contains(&id2));
306 assert!(ids.contains(&id3));
307 }
308
309 #[test]
310 fn test_remove_matching() {
311 let store = SessionStore::new();
312 store.create().unwrap();
313 store.create().unwrap();
314
315 let ids = store.list_ids().unwrap();
317 store
318 .update(&ids[0], |s| s.state = SessionState::Terminated)
319 .unwrap();
320
321 let removed = store
323 .remove_matching(|s| s.state == SessionState::Terminated)
324 .unwrap();
325
326 assert_eq!(removed, 1);
327 assert_eq!(store.count(), 1);
328 }
329
330 #[test]
335 fn an_idle_session_is_swept() {
336 let store = SessionStore::new();
337 let id = store.create().unwrap();
338 store
339 .update(&id, |s| {
340 let _ = s.state.transition_to(SessionState::Idle);
341 })
342 .unwrap();
343
344 let swept = store.sweep_idle(std::time::Duration::ZERO).unwrap();
346
347 assert_eq!(swept, vec![id], "the idle session must be the one reported");
348 assert_eq!(store.count(), 0);
349 }
350
351 #[test]
360 fn a_session_running_a_command_is_never_swept() {
361 let store = SessionStore::new();
362 let running = store.create().unwrap();
363 let abandoned = store.create().unwrap();
364 store
365 .update(&running, |s| {
366 let _ = s.state.transition_to(SessionState::Active);
367 })
368 .unwrap();
369 store
370 .update(&abandoned, |s| {
371 let _ = s.state.transition_to(SessionState::Idle);
372 })
373 .unwrap();
374
375 let swept = store.sweep_idle(std::time::Duration::ZERO).unwrap();
376
377 assert_eq!(
378 swept,
379 vec![abandoned],
380 "only the abandoned session may be swept"
381 );
382 assert!(
383 store.contains(&running).unwrap(),
384 "a session with a command running in it must survive a sweep"
385 );
386 }
387
388 #[test]
391 fn a_session_within_the_ttl_stays() {
392 let store = SessionStore::new();
393 let id = store.create().unwrap();
394
395 let swept = store
396 .sweep_idle(std::time::Duration::from_secs(3600))
397 .unwrap();
398
399 assert!(swept.is_empty(), "nothing has been idle for an hour yet");
400 assert!(store.contains(&id).unwrap());
401 }
402
403 #[test]
404 fn test_concurrent_access() {
405 use std::sync::Arc;
406 use std::thread;
407
408 let store = Arc::new(SessionStore::new());
409 let mut handles = vec![];
410
411 for _ in 0..100 {
413 let store = Arc::clone(&store);
414 handles.push(thread::spawn(move || store.create().unwrap()));
415 }
416
417 let ids: Vec<SessionId> = handles.into_iter().map(|h| h.join().unwrap()).collect();
418
419 let unique: std::collections::HashSet<_> = ids.iter().collect();
421 assert_eq!(unique.len(), 100);
422
423 assert_eq!(store.count(), 100);
425 }
426}