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>> {
192 let mut sessions = self
193 .sessions
194 .write()
195 .map_err(|_| ShellTunnelError::LockPoisoned)?;
196
197 let expired: Vec<SessionId> = sessions
198 .values()
199 .filter(|s| s.state != SessionState::Active && s.idle_duration() > ttl)
200 .map(|s| s.id)
201 .collect();
202
203 for id in &expired {
204 sessions.remove(id);
205 }
206 Ok(expired)
207 }
208}
209
210impl Default for SessionStore {
211 fn default() -> Self {
212 Self::new()
213 }
214}
215
216#[cfg(test)]
217mod tests {
218 use super::*;
219
220 #[test]
221 fn test_create_session() {
222 let store = SessionStore::new();
223 let id = store.create().unwrap();
224
225 assert!(store.contains(&id).unwrap());
226 assert_eq!(store.count(), 1);
227 }
228
229 #[test]
230 fn test_get_session() {
231 let store = SessionStore::new();
232 let id = store.create().unwrap();
233
234 let session = store.get(&id).unwrap().unwrap();
235 assert_eq!(session.id, id);
236 assert_eq!(session.state, SessionState::Created);
237 }
238
239 #[test]
240 fn test_get_nonexistent() {
241 let store = SessionStore::new();
242 let fake_id = SessionId::from_raw(999999);
243
244 let result = store.get(&fake_id).unwrap();
245 assert!(result.is_none());
246 }
247
248 #[test]
249 fn test_update_session() {
250 let store = SessionStore::new();
251 let id = store.create().unwrap();
252
253 store
254 .update(&id, |s| {
255 s.state = SessionState::Active;
256 })
257 .unwrap();
258
259 let session = store.get(&id).unwrap().unwrap();
260 assert_eq!(session.state, SessionState::Active);
261 }
262
263 #[test]
264 fn test_update_nonexistent() {
265 let store = SessionStore::new();
266 let fake_id = SessionId::from_raw(999999);
267
268 let result = store.update(&fake_id, |_| {});
269 assert!(result.is_err());
270 }
271
272 #[test]
273 fn test_remove_session() {
274 let store = SessionStore::new();
275 let id = store.create().unwrap();
276
277 let removed = store.remove(&id).unwrap();
278 assert!(removed.is_some());
279 assert_eq!(removed.unwrap().id, id);
280
281 assert!(!store.contains(&id).unwrap());
282 assert_eq!(store.count(), 0);
283 }
284
285 #[test]
286 fn test_list_ids() {
287 let store = SessionStore::new();
288 let id1 = store.create().unwrap();
289 let id2 = store.create().unwrap();
290 let id3 = store.create().unwrap();
291
292 let ids = store.list_ids().unwrap();
293 assert_eq!(ids.len(), 3);
294 assert!(ids.contains(&id1));
295 assert!(ids.contains(&id2));
296 assert!(ids.contains(&id3));
297 }
298
299 #[test]
300 fn test_remove_matching() {
301 let store = SessionStore::new();
302 store.create().unwrap();
303 store.create().unwrap();
304
305 let ids = store.list_ids().unwrap();
307 store
308 .update(&ids[0], |s| s.state = SessionState::Terminated)
309 .unwrap();
310
311 let removed = store
313 .remove_matching(|s| s.state == SessionState::Terminated)
314 .unwrap();
315
316 assert_eq!(removed, 1);
317 assert_eq!(store.count(), 1);
318 }
319
320 #[test]
325 fn an_idle_session_is_swept() {
326 let store = SessionStore::new();
327 let id = store.create().unwrap();
328 store
329 .update(&id, |s| {
330 let _ = s.state.transition_to(SessionState::Idle);
331 })
332 .unwrap();
333
334 let swept = store.sweep_idle(std::time::Duration::ZERO).unwrap();
336
337 assert_eq!(swept, vec![id], "the idle session must be the one reported");
338 assert_eq!(store.count(), 0);
339 }
340
341 #[test]
350 fn a_session_running_a_command_is_never_swept() {
351 let store = SessionStore::new();
352 let running = store.create().unwrap();
353 let abandoned = store.create().unwrap();
354 store
355 .update(&running, |s| {
356 let _ = s.state.transition_to(SessionState::Active);
357 })
358 .unwrap();
359 store
360 .update(&abandoned, |s| {
361 let _ = s.state.transition_to(SessionState::Idle);
362 })
363 .unwrap();
364
365 let swept = store.sweep_idle(std::time::Duration::ZERO).unwrap();
366
367 assert_eq!(
368 swept,
369 vec![abandoned],
370 "only the abandoned session may be swept"
371 );
372 assert!(
373 store.contains(&running).unwrap(),
374 "a session with a command running in it must survive a sweep"
375 );
376 }
377
378 #[test]
381 fn a_session_within_the_ttl_stays() {
382 let store = SessionStore::new();
383 let id = store.create().unwrap();
384
385 let swept = store
386 .sweep_idle(std::time::Duration::from_secs(3600))
387 .unwrap();
388
389 assert!(swept.is_empty(), "nothing has been idle for an hour yet");
390 assert!(store.contains(&id).unwrap());
391 }
392
393 #[test]
394 fn test_concurrent_access() {
395 use std::sync::Arc;
396 use std::thread;
397
398 let store = Arc::new(SessionStore::new());
399 let mut handles = vec![];
400
401 for _ in 0..100 {
403 let store = Arc::clone(&store);
404 handles.push(thread::spawn(move || store.create().unwrap()));
405 }
406
407 let ids: Vec<SessionId> = handles.into_iter().map(|h| h.join().unwrap()).collect();
408
409 let unique: std::collections::HashSet<_> = ids.iter().collect();
411 assert_eq!(unique.len(), 100);
412
413 assert_eq!(store.count(), 100);
415 }
416}