1use std::sync::Arc;
10use std::sync::RwLock;
11use std::sync::atomic::{AtomicBool, 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 handshake_started: Arc<AtomicBool>,
72 extensions: Arc<RwLock<Extensions>>,
73}
74
75impl Default for SessionState {
76 fn default() -> Self {
77 Self::new()
78 }
79}
80
81impl SessionState {
82 pub fn new() -> Self {
84 Self {
85 phase: Arc::new(AtomicU8::new(SessionPhase::Uninitialized as u8)),
86 handshake_started: Arc::new(AtomicBool::new(false)),
87 extensions: Arc::new(RwLock::new(Extensions::new())),
88 }
89 }
90
91 pub fn insert<T: Send + Sync + Clone + 'static>(&self, val: T) {
106 if let Ok(mut ext) = self.extensions.write() {
107 ext.insert(val);
108 }
109 }
110
111 pub fn get<T: Send + Sync + Clone + 'static>(&self) -> Option<T> {
127 self.extensions
128 .read()
129 .ok()
130 .and_then(|ext| ext.get::<T>().cloned())
131 }
132
133 pub fn phase(&self) -> SessionPhase {
135 SessionPhase::from(self.phase.load(Ordering::Acquire))
136 }
137
138 pub fn is_initialized(&self) -> bool {
140 self.phase() == SessionPhase::Initialized
141 }
142
143 pub fn mark_handshake_started(&self) {
155 self.handshake_started.store(true, Ordering::Release);
156 }
157
158 pub fn handshake_started(&self) -> bool {
163 self.handshake_started.load(Ordering::Acquire)
164 }
165
166 pub fn mark_initializing(&self) -> bool {
173 self.mark_handshake_started();
174 self.phase
175 .compare_exchange(
176 SessionPhase::Uninitialized as u8,
177 SessionPhase::Initializing as u8,
178 Ordering::AcqRel,
179 Ordering::Acquire,
180 )
181 .is_ok()
182 }
183
184 pub fn mark_initialized(&self) -> bool {
202 if self
204 .phase
205 .compare_exchange(
206 SessionPhase::Initializing as u8,
207 SessionPhase::Initialized as u8,
208 Ordering::AcqRel,
209 Ordering::Acquire,
210 )
211 .is_ok()
212 {
213 return true;
214 }
215
216 if !self.handshake_started() {
220 return false;
221 }
222
223 self.phase
224 .compare_exchange(
225 SessionPhase::Uninitialized as u8,
226 SessionPhase::Initialized as u8,
227 Ordering::AcqRel,
228 Ordering::Acquire,
229 )
230 .is_ok()
231 }
232
233 pub fn mark_preinitialized(&self) -> bool {
246 self.mark_handshake_started();
247 self.phase
248 .swap(SessionPhase::Initialized as u8, Ordering::AcqRel)
249 != SessionPhase::Initialized as u8
250 }
251
252 pub fn is_request_allowed(&self, method: &str) -> bool {
257 match self.phase() {
258 SessionPhase::Uninitialized => {
259 matches!(method, "initialize" | "ping" | "server/discover")
261 }
262 SessionPhase::Initializing | SessionPhase::Initialized => true,
263 }
264 }
265}
266
267#[cfg(test)]
268mod tests {
269 use super::*;
270 use proptest::prelude::*;
271
272 #[derive(Clone, Debug)]
273 enum LifecycleOperation {
274 MarkInitializing,
275 MarkInitialized,
276 MarkHandshakeStarted,
277 MarkPreinitialized,
278 CheckRequest(String),
279 }
280
281 fn lifecycle_operation() -> impl Strategy<Value = LifecycleOperation> {
282 prop_oneof![
283 Just(LifecycleOperation::MarkInitializing),
284 Just(LifecycleOperation::MarkInitialized),
285 Just(LifecycleOperation::MarkHandshakeStarted),
286 Just(LifecycleOperation::MarkPreinitialized),
287 prop_oneof![
288 Just("initialize".to_string()),
289 Just("ping".to_string()),
290 Just("server/discover".to_string()),
291 Just("tools/list".to_string()),
292 "[a-z/_.-]{0,64}",
293 ]
294 .prop_map(LifecycleOperation::CheckRequest),
295 ]
296 }
297
298 proptest! {
299 #![proptest_config(ProptestConfig::with_cases(512))]
300
301 #[test]
304 fn lifecycle_matches_model(
305 operations in prop::collection::vec(lifecycle_operation(), 0..256)
306 ) {
307 let session = SessionState::new();
308 let mut expected_phase = SessionPhase::Uninitialized;
309 let mut expected_handshake = false;
310
311 for operation in operations {
312 match operation {
313 LifecycleOperation::MarkInitializing => {
314 let expected_success = expected_phase == SessionPhase::Uninitialized;
315 prop_assert_eq!(session.mark_initializing(), expected_success);
316 expected_handshake = true;
317 if expected_success {
318 expected_phase = SessionPhase::Initializing;
319 }
320 }
321 LifecycleOperation::MarkInitialized => {
322 let expected_success = match expected_phase {
325 SessionPhase::Initializing => true,
326 SessionPhase::Uninitialized => expected_handshake,
327 SessionPhase::Initialized => false,
328 };
329 prop_assert_eq!(session.mark_initialized(), expected_success);
330 if expected_success {
331 expected_phase = SessionPhase::Initialized;
332 }
333 }
334 LifecycleOperation::MarkHandshakeStarted => {
335 session.mark_handshake_started();
336 expected_handshake = true;
337 }
338 LifecycleOperation::MarkPreinitialized => {
339 let expected_success = expected_phase != SessionPhase::Initialized;
340 prop_assert_eq!(session.mark_preinitialized(), expected_success);
341 expected_handshake = true;
342 expected_phase = SessionPhase::Initialized;
343 }
344 LifecycleOperation::CheckRequest(method) => {
345 let expected_allowed = expected_phase != SessionPhase::Uninitialized
346 || matches!(method.as_str(), "initialize" | "ping" | "server/discover");
347 prop_assert_eq!(
348 session.is_request_allowed(&method),
349 expected_allowed,
350 "phase={:?}, method={:?}",
351 expected_phase,
352 method
353 );
354 }
355 }
356 prop_assert_eq!(session.phase(), expected_phase);
357 prop_assert_eq!(session.handshake_started(), expected_handshake);
358 prop_assert_eq!(
359 session.is_initialized(),
360 expected_phase == SessionPhase::Initialized
361 );
362 prop_assert!(expected_phase != SessionPhase::Initialized || expected_handshake);
365 }
366 }
367 }
368
369 #[test]
370 fn test_session_lifecycle() {
371 let session = SessionState::new();
372
373 assert_eq!(session.phase(), SessionPhase::Uninitialized);
375 assert!(!session.is_initialized());
376
377 assert!(session.is_request_allowed("initialize"));
379 assert!(session.is_request_allowed("ping"));
380 assert!(!session.is_request_allowed("tools/list"));
381
382 assert!(session.mark_initializing());
384 assert_eq!(session.phase(), SessionPhase::Initializing);
385 assert!(!session.is_initialized());
386
387 assert!(!session.mark_initializing());
389
390 assert!(session.is_request_allowed("tools/list"));
392
393 assert!(session.mark_initialized());
395 assert_eq!(session.phase(), SessionPhase::Initialized);
396 assert!(session.is_initialized());
397
398 assert!(!session.mark_initialized());
400 }
401
402 #[test]
403 fn test_session_clone_shares_state() {
404 let session1 = SessionState::new();
405 let session2 = session1.clone();
406
407 session1.mark_initializing();
408 assert_eq!(session2.phase(), SessionPhase::Initializing);
409
410 session2.mark_initialized();
411 assert_eq!(session1.phase(), SessionPhase::Initialized);
412 }
413
414 #[test]
415 fn test_session_extensions_insert_and_get() {
416 let session = SessionState::new();
417
418 session.insert(42u32);
420 assert_eq!(session.get::<u32>(), Some(42));
421
422 assert_eq!(session.get::<String>(), None);
424 }
425
426 #[test]
427 fn test_session_extensions_overwrite() {
428 let session = SessionState::new();
429
430 session.insert(42u32);
431 assert_eq!(session.get::<u32>(), Some(42));
432
433 session.insert(100u32);
435 assert_eq!(session.get::<u32>(), Some(100));
436 }
437
438 #[test]
439 fn test_session_extensions_multiple_types() {
440 let session = SessionState::new();
441
442 session.insert(42u32);
443 session.insert("hello".to_string());
444 session.insert(true);
445
446 assert_eq!(session.get::<u32>(), Some(42));
447 assert_eq!(session.get::<String>(), Some("hello".to_string()));
448 assert_eq!(session.get::<bool>(), Some(true));
449 }
450
451 #[test]
452 fn test_session_extensions_shared_across_clones() {
453 let session1 = SessionState::new();
454 let session2 = session1.clone();
455
456 session1.insert(42u32);
458
459 assert_eq!(session2.get::<u32>(), Some(42));
461
462 session2.insert("world".to_string());
464
465 assert_eq!(session1.get::<String>(), Some("world".to_string()));
467 }
468
469 #[test]
474 fn test_mark_initialized_from_uninitialized_after_handshake_started() {
475 let session = SessionState::new();
476 session.mark_handshake_started();
477
478 assert_eq!(session.phase(), SessionPhase::Uninitialized);
479 assert!(session.mark_initialized());
480 assert_eq!(session.phase(), SessionPhase::Initialized);
481 assert!(session.is_initialized());
482
483 assert!(session.is_request_allowed("tools/list"));
485 assert!(session.is_request_allowed("ping"));
486 }
487
488 #[test]
492 fn test_mark_initialized_from_uninitialized_without_handshake_is_refused() {
493 let session = SessionState::new();
494
495 assert!(!session.handshake_started());
496 assert!(!session.mark_initialized());
497 assert_eq!(session.phase(), SessionPhase::Uninitialized);
498 assert!(!session.is_initialized());
499
500 assert!(!session.is_request_allowed("tools/list"));
502 assert!(session.is_request_allowed("initialize"));
503 assert!(session.is_request_allowed("ping"));
504
505 assert!(!session.mark_initialized());
507 assert!(!session.mark_initialized());
508 assert_eq!(session.phase(), SessionPhase::Uninitialized);
509 }
510
511 #[test]
513 fn test_handshake_still_works_after_a_refused_notification() {
514 let session = SessionState::new();
515
516 assert!(!session.mark_initialized());
517 assert!(session.mark_initializing());
518 assert_eq!(session.phase(), SessionPhase::Initializing);
519 assert!(session.mark_initialized());
520 assert!(session.is_initialized());
521 }
522
523 #[test]
525 fn test_mark_preinitialized_skips_the_handshake() {
526 let session = SessionState::new();
527 assert!(session.mark_preinitialized());
528 assert_eq!(session.phase(), SessionPhase::Initialized);
529 assert!(session.handshake_started());
530 assert!(session.is_request_allowed("tools/list"));
531
532 assert!(!session.mark_preinitialized());
534
535 let mid = SessionState::new();
537 mid.mark_initializing();
538 assert!(mid.mark_preinitialized());
539 assert_eq!(mid.phase(), SessionPhase::Initialized);
540 }
541
542 #[test]
543 fn test_handshake_flag_is_shared_across_clones() {
544 let session1 = SessionState::new();
545 let session2 = session1.clone();
546
547 assert!(!session2.handshake_started());
548 session1.mark_handshake_started();
549 assert!(session2.handshake_started());
550
551 assert!(session2.mark_initialized());
553 assert!(session1.is_initialized());
554 }
555
556 #[test]
557 fn test_mark_initialized_idempotent_when_already_initialized() {
558 let session = SessionState::new();
559
560 session.mark_initializing();
562 session.mark_initialized();
563 assert_eq!(session.phase(), SessionPhase::Initialized);
564
565 assert!(!session.mark_initialized());
567 assert_eq!(session.phase(), SessionPhase::Initialized);
568 }
569
570 #[test]
571 fn test_session_extensions_custom_type() {
572 #[derive(Debug, Clone, PartialEq)]
573 struct UserClaims {
574 user_id: String,
575 role: String,
576 }
577
578 let session = SessionState::new();
579
580 session.insert(UserClaims {
581 user_id: "user123".to_string(),
582 role: "admin".to_string(),
583 });
584
585 let claims = session.get::<UserClaims>();
586 assert!(claims.is_some());
587 let claims = claims.unwrap();
588 assert_eq!(claims.user_id, "user123");
589 assert_eq!(claims.role, "admin");
590 }
591}