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
168impl Default for SessionStore {
169 fn default() -> Self {
170 Self::new()
171 }
172}
173
174#[cfg(test)]
175mod tests {
176 use super::*;
177
178 #[test]
179 fn test_create_session() {
180 let store = SessionStore::new();
181 let id = store.create().unwrap();
182
183 assert!(store.contains(&id).unwrap());
184 assert_eq!(store.count(), 1);
185 }
186
187 #[test]
188 fn test_get_session() {
189 let store = SessionStore::new();
190 let id = store.create().unwrap();
191
192 let session = store.get(&id).unwrap().unwrap();
193 assert_eq!(session.id, id);
194 assert_eq!(session.state, SessionState::Created);
195 }
196
197 #[test]
198 fn test_get_nonexistent() {
199 let store = SessionStore::new();
200 let fake_id = SessionId::from_raw(999999);
201
202 let result = store.get(&fake_id).unwrap();
203 assert!(result.is_none());
204 }
205
206 #[test]
207 fn test_update_session() {
208 let store = SessionStore::new();
209 let id = store.create().unwrap();
210
211 store
212 .update(&id, |s| {
213 s.state = SessionState::Active;
214 })
215 .unwrap();
216
217 let session = store.get(&id).unwrap().unwrap();
218 assert_eq!(session.state, SessionState::Active);
219 }
220
221 #[test]
222 fn test_update_nonexistent() {
223 let store = SessionStore::new();
224 let fake_id = SessionId::from_raw(999999);
225
226 let result = store.update(&fake_id, |_| {});
227 assert!(result.is_err());
228 }
229
230 #[test]
231 fn test_remove_session() {
232 let store = SessionStore::new();
233 let id = store.create().unwrap();
234
235 let removed = store.remove(&id).unwrap();
236 assert!(removed.is_some());
237 assert_eq!(removed.unwrap().id, id);
238
239 assert!(!store.contains(&id).unwrap());
240 assert_eq!(store.count(), 0);
241 }
242
243 #[test]
244 fn test_list_ids() {
245 let store = SessionStore::new();
246 let id1 = store.create().unwrap();
247 let id2 = store.create().unwrap();
248 let id3 = store.create().unwrap();
249
250 let ids = store.list_ids().unwrap();
251 assert_eq!(ids.len(), 3);
252 assert!(ids.contains(&id1));
253 assert!(ids.contains(&id2));
254 assert!(ids.contains(&id3));
255 }
256
257 #[test]
258 fn test_remove_matching() {
259 let store = SessionStore::new();
260 store.create().unwrap();
261 store.create().unwrap();
262
263 let ids = store.list_ids().unwrap();
265 store
266 .update(&ids[0], |s| s.state = SessionState::Terminated)
267 .unwrap();
268
269 let removed = store
271 .remove_matching(|s| s.state == SessionState::Terminated)
272 .unwrap();
273
274 assert_eq!(removed, 1);
275 assert_eq!(store.count(), 1);
276 }
277
278 #[test]
279 fn test_concurrent_access() {
280 use std::sync::Arc;
281 use std::thread;
282
283 let store = Arc::new(SessionStore::new());
284 let mut handles = vec![];
285
286 for _ in 0..100 {
288 let store = Arc::clone(&store);
289 handles.push(thread::spawn(move || store.create().unwrap()));
290 }
291
292 let ids: Vec<SessionId> = handles.into_iter().map(|h| h.join().unwrap()).collect();
293
294 let unique: std::collections::HashSet<_> = ids.iter().collect();
296 assert_eq!(unique.len(), 100);
297
298 assert_eq!(store.count(), 100);
300 }
301}