1use std::sync::Arc;
10use std::sync::RwLock;
11use std::sync::atomic::{AtomicU8, Ordering};
12
13use crate::router::Extensions;
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17#[repr(u8)]
18#[non_exhaustive]
19pub enum SessionPhase {
20 Uninitialized = 0,
22 Initializing = 1,
24 Initialized = 2,
26}
27
28impl From<u8> for SessionPhase {
29 fn from(value: u8) -> Self {
30 match value {
31 0 => SessionPhase::Uninitialized,
32 1 => SessionPhase::Initializing,
33 2 => SessionPhase::Initialized,
34 _ => SessionPhase::Uninitialized,
35 }
36 }
37}
38
39#[derive(Clone)]
69pub struct SessionState {
70 phase: Arc<AtomicU8>,
71 extensions: Arc<RwLock<Extensions>>,
72}
73
74impl Default for SessionState {
75 fn default() -> Self {
76 Self::new()
77 }
78}
79
80impl SessionState {
81 pub fn new() -> Self {
83 Self {
84 phase: Arc::new(AtomicU8::new(SessionPhase::Uninitialized as u8)),
85 extensions: Arc::new(RwLock::new(Extensions::new())),
86 }
87 }
88
89 pub fn insert<T: Send + Sync + Clone + 'static>(&self, val: T) {
104 if let Ok(mut ext) = self.extensions.write() {
105 ext.insert(val);
106 }
107 }
108
109 pub fn get<T: Send + Sync + Clone + 'static>(&self) -> Option<T> {
125 self.extensions
126 .read()
127 .ok()
128 .and_then(|ext| ext.get::<T>().cloned())
129 }
130
131 pub fn phase(&self) -> SessionPhase {
133 SessionPhase::from(self.phase.load(Ordering::Acquire))
134 }
135
136 pub fn is_initialized(&self) -> bool {
138 self.phase() == SessionPhase::Initialized
139 }
140
141 pub fn mark_initializing(&self) -> bool {
145 self.phase
146 .compare_exchange(
147 SessionPhase::Uninitialized as u8,
148 SessionPhase::Initializing as u8,
149 Ordering::AcqRel,
150 Ordering::Acquire,
151 )
152 .is_ok()
153 }
154
155 pub fn mark_initialized(&self) -> bool {
166 if self
168 .phase
169 .compare_exchange(
170 SessionPhase::Initializing as u8,
171 SessionPhase::Initialized as u8,
172 Ordering::AcqRel,
173 Ordering::Acquire,
174 )
175 .is_ok()
176 {
177 return true;
178 }
179
180 self.phase
184 .compare_exchange(
185 SessionPhase::Uninitialized as u8,
186 SessionPhase::Initialized as u8,
187 Ordering::AcqRel,
188 Ordering::Acquire,
189 )
190 .is_ok()
191 }
192
193 pub fn is_request_allowed(&self, method: &str) -> bool {
198 match self.phase() {
199 SessionPhase::Uninitialized => {
200 matches!(method, "initialize" | "ping" | "server/discover")
202 }
203 SessionPhase::Initializing | SessionPhase::Initialized => true,
204 }
205 }
206}
207
208#[cfg(test)]
209mod tests {
210 use super::*;
211 use proptest::prelude::*;
212
213 #[derive(Clone, Debug)]
214 enum LifecycleOperation {
215 MarkInitializing,
216 MarkInitialized,
217 CheckRequest(String),
218 }
219
220 fn lifecycle_operation() -> impl Strategy<Value = LifecycleOperation> {
221 prop_oneof![
222 Just(LifecycleOperation::MarkInitializing),
223 Just(LifecycleOperation::MarkInitialized),
224 prop_oneof![
225 Just("initialize".to_string()),
226 Just("ping".to_string()),
227 Just("server/discover".to_string()),
228 Just("tools/list".to_string()),
229 "[a-z/_.-]{0,64}",
230 ]
231 .prop_map(LifecycleOperation::CheckRequest),
232 ]
233 }
234
235 proptest! {
236 #![proptest_config(ProptestConfig::with_cases(512))]
237
238 #[test]
241 fn lifecycle_matches_model(
242 operations in prop::collection::vec(lifecycle_operation(), 0..256)
243 ) {
244 let session = SessionState::new();
245 let mut expected_phase = SessionPhase::Uninitialized;
246
247 for operation in operations {
248 match operation {
249 LifecycleOperation::MarkInitializing => {
250 let expected_success = expected_phase == SessionPhase::Uninitialized;
251 prop_assert_eq!(session.mark_initializing(), expected_success);
252 if expected_success {
253 expected_phase = SessionPhase::Initializing;
254 }
255 }
256 LifecycleOperation::MarkInitialized => {
257 let expected_success = expected_phase != SessionPhase::Initialized;
258 prop_assert_eq!(session.mark_initialized(), expected_success);
259 if expected_success {
260 expected_phase = SessionPhase::Initialized;
261 }
262 }
263 LifecycleOperation::CheckRequest(method) => {
264 let expected_allowed = expected_phase != SessionPhase::Uninitialized
265 || matches!(method.as_str(), "initialize" | "ping" | "server/discover");
266 prop_assert_eq!(
267 session.is_request_allowed(&method),
268 expected_allowed,
269 "phase={:?}, method={:?}",
270 expected_phase,
271 method
272 );
273 }
274 }
275 prop_assert_eq!(session.phase(), expected_phase);
276 prop_assert_eq!(
277 session.is_initialized(),
278 expected_phase == SessionPhase::Initialized
279 );
280 }
281 }
282 }
283
284 #[test]
285 fn test_session_lifecycle() {
286 let session = SessionState::new();
287
288 assert_eq!(session.phase(), SessionPhase::Uninitialized);
290 assert!(!session.is_initialized());
291
292 assert!(session.is_request_allowed("initialize"));
294 assert!(session.is_request_allowed("ping"));
295 assert!(!session.is_request_allowed("tools/list"));
296
297 assert!(session.mark_initializing());
299 assert_eq!(session.phase(), SessionPhase::Initializing);
300 assert!(!session.is_initialized());
301
302 assert!(!session.mark_initializing());
304
305 assert!(session.is_request_allowed("tools/list"));
307
308 assert!(session.mark_initialized());
310 assert_eq!(session.phase(), SessionPhase::Initialized);
311 assert!(session.is_initialized());
312
313 assert!(!session.mark_initialized());
315 }
316
317 #[test]
318 fn test_session_clone_shares_state() {
319 let session1 = SessionState::new();
320 let session2 = session1.clone();
321
322 session1.mark_initializing();
323 assert_eq!(session2.phase(), SessionPhase::Initializing);
324
325 session2.mark_initialized();
326 assert_eq!(session1.phase(), SessionPhase::Initialized);
327 }
328
329 #[test]
330 fn test_session_extensions_insert_and_get() {
331 let session = SessionState::new();
332
333 session.insert(42u32);
335 assert_eq!(session.get::<u32>(), Some(42));
336
337 assert_eq!(session.get::<String>(), None);
339 }
340
341 #[test]
342 fn test_session_extensions_overwrite() {
343 let session = SessionState::new();
344
345 session.insert(42u32);
346 assert_eq!(session.get::<u32>(), Some(42));
347
348 session.insert(100u32);
350 assert_eq!(session.get::<u32>(), Some(100));
351 }
352
353 #[test]
354 fn test_session_extensions_multiple_types() {
355 let session = SessionState::new();
356
357 session.insert(42u32);
358 session.insert("hello".to_string());
359 session.insert(true);
360
361 assert_eq!(session.get::<u32>(), Some(42));
362 assert_eq!(session.get::<String>(), Some("hello".to_string()));
363 assert_eq!(session.get::<bool>(), Some(true));
364 }
365
366 #[test]
367 fn test_session_extensions_shared_across_clones() {
368 let session1 = SessionState::new();
369 let session2 = session1.clone();
370
371 session1.insert(42u32);
373
374 assert_eq!(session2.get::<u32>(), Some(42));
376
377 session2.insert("world".to_string());
379
380 assert_eq!(session1.get::<String>(), Some("world".to_string()));
382 }
383
384 #[test]
385 fn test_mark_initialized_from_uninitialized() {
386 let session = SessionState::new();
387
388 assert_eq!(session.phase(), SessionPhase::Uninitialized);
392 assert!(session.mark_initialized());
393 assert_eq!(session.phase(), SessionPhase::Initialized);
394 assert!(session.is_initialized());
395
396 assert!(session.is_request_allowed("tools/list"));
398 assert!(session.is_request_allowed("ping"));
399 }
400
401 #[test]
402 fn test_mark_initialized_idempotent_when_already_initialized() {
403 let session = SessionState::new();
404
405 session.mark_initializing();
407 session.mark_initialized();
408 assert_eq!(session.phase(), SessionPhase::Initialized);
409
410 assert!(!session.mark_initialized());
412 assert_eq!(session.phase(), SessionPhase::Initialized);
413 }
414
415 #[test]
416 fn test_session_extensions_custom_type() {
417 #[derive(Debug, Clone, PartialEq)]
418 struct UserClaims {
419 user_id: String,
420 role: String,
421 }
422
423 let session = SessionState::new();
424
425 session.insert(UserClaims {
426 user_id: "user123".to_string(),
427 role: "admin".to_string(),
428 });
429
430 let claims = session.get::<UserClaims>();
431 assert!(claims.is_some());
432 let claims = claims.unwrap();
433 assert_eq!(claims.user_id, "user123");
434 assert_eq!(claims.role, "admin");
435 }
436}