Skip to main content

tower_mcp/
session.rs

1//! MCP session state management
2//!
3//! Tracks the lifecycle state of an MCP connection as per the specification.
4//! The session progresses through phases: Uninitialized -> Initializing -> Initialized.
5//!
6//! Sessions also support type-safe extensions for storing arbitrary data like
7//! authentication claims, user roles, or other session-scoped state.
8
9use std::sync::Arc;
10use std::sync::RwLock;
11use std::sync::atomic::{AtomicU8, Ordering};
12
13use crate::router::Extensions;
14
15/// Session lifecycle phase
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17#[repr(u8)]
18#[non_exhaustive]
19pub enum SessionPhase {
20    /// Initial state - only `initialize` and `ping` requests are valid
21    Uninitialized = 0,
22    /// Server has responded to `initialize`, waiting for `initialized` notification
23    Initializing = 1,
24    /// `initialized` notification received, normal operation
25    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/// Shared session state that can be cloned across requests.
40///
41/// Uses atomic operations for thread-safe state transitions. Includes a type-safe
42/// extensions map for storing session-scoped data like authentication claims.
43///
44/// # Example
45///
46/// ```rust
47/// use tower_mcp::SessionState;
48///
49/// #[derive(Debug, Clone)]
50/// struct UserClaims {
51///     user_id: String,
52///     role: String,
53/// }
54///
55/// let session = SessionState::new();
56///
57/// // Store auth claims in the session
58/// session.insert(UserClaims {
59///     user_id: "user123".to_string(),
60///     role: "admin".to_string(),
61/// });
62///
63/// // Retrieve claims later
64/// if let Some(claims) = session.get::<UserClaims>() {
65///     assert_eq!(claims.role, "admin");
66/// }
67/// ```
68#[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    /// Create a new session in the Uninitialized phase
82    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    /// Insert a value into the session extensions.
90    ///
91    /// This is typically used by auth middleware to store claims that can
92    /// be checked by capability filters.
93    ///
94    /// # Example
95    ///
96    /// ```rust
97    /// use tower_mcp::SessionState;
98    ///
99    /// let session = SessionState::new();
100    /// session.insert(42u32);
101    /// assert_eq!(session.get::<u32>(), Some(42));
102    /// ```
103    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    /// Get a cloned value from the session extensions.
110    ///
111    /// Returns `None` if no value of the given type has been inserted or if
112    /// the lock cannot be acquired.
113    ///
114    /// # Example
115    ///
116    /// ```rust
117    /// use tower_mcp::SessionState;
118    ///
119    /// let session = SessionState::new();
120    /// session.insert("hello".to_string());
121    /// assert_eq!(session.get::<String>(), Some("hello".to_string()));
122    /// assert_eq!(session.get::<u32>(), None);
123    /// ```
124    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    /// Get the current session phase
132    pub fn phase(&self) -> SessionPhase {
133        SessionPhase::from(self.phase.load(Ordering::Acquire))
134    }
135
136    /// Check if the session is initialized (operation phase)
137    pub fn is_initialized(&self) -> bool {
138        self.phase() == SessionPhase::Initialized
139    }
140
141    /// Transition from Uninitialized to Initializing.
142    /// Called after responding to an `initialize` request.
143    /// Returns true if the transition was successful.
144    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    /// Transition to Initialized phase.
156    /// Called when receiving an `initialized` notification.
157    ///
158    /// Accepts transitions from both `Initializing` and `Uninitialized` states.
159    /// The `Uninitialized → Initialized` path handles a race condition in HTTP
160    /// transports where the client sends the `initialized` notification before
161    /// the server has finished processing the `initialize` request (the session
162    /// is stored in `Uninitialized` state before the request is dispatched).
163    ///
164    /// Returns true if the transition was successful.
165    pub fn mark_initialized(&self) -> bool {
166        // Try the expected path first: Initializing → Initialized
167        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        // Handle the race: Uninitialized → Initialized
181        // This occurs when the initialized notification arrives before
182        // the initialize request has been fully processed.
183        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    /// Check if a request method is allowed in the current phase.
194    /// Per spec:
195    /// - Before initialization: only `initialize` and `ping` are valid
196    /// - During all phases: `ping` is always valid
197    pub fn is_request_allowed(&self, method: &str) -> bool {
198        match self.phase() {
199            SessionPhase::Uninitialized => {
200                // server/discover (SEP-1442) is allowed before initialization
201                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        /// Model random lifecycle sequences and assert that the atomic state
239        /// machine never permits an illegal transition or pre-init request.
240        #[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        // Initial state
289        assert_eq!(session.phase(), SessionPhase::Uninitialized);
290        assert!(!session.is_initialized());
291
292        // Only initialize and ping allowed
293        assert!(session.is_request_allowed("initialize"));
294        assert!(session.is_request_allowed("ping"));
295        assert!(!session.is_request_allowed("tools/list"));
296
297        // Transition to initializing
298        assert!(session.mark_initializing());
299        assert_eq!(session.phase(), SessionPhase::Initializing);
300        assert!(!session.is_initialized());
301
302        // Can't mark initializing again
303        assert!(!session.mark_initializing());
304
305        // All requests allowed during initializing
306        assert!(session.is_request_allowed("tools/list"));
307
308        // Transition to initialized
309        assert!(session.mark_initialized());
310        assert_eq!(session.phase(), SessionPhase::Initialized);
311        assert!(session.is_initialized());
312
313        // Can't mark initialized again
314        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        // Insert and retrieve a value
334        session.insert(42u32);
335        assert_eq!(session.get::<u32>(), Some(42));
336
337        // Different type returns None
338        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        // Overwrite with new value
349        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        // Insert in one clone
372        session1.insert(42u32);
373
374        // Should be visible in the other
375        assert_eq!(session2.get::<u32>(), Some(42));
376
377        // Insert in the second clone
378        session2.insert("world".to_string());
379
380        // Should be visible in the first
381        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        // Start in Uninitialized, skip straight to Initialized
389        // This handles the race where `initialized` notification arrives
390        // before the `initialize` request is fully processed.
391        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        // All requests allowed
397        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        // Normal lifecycle
406        session.mark_initializing();
407        session.mark_initialized();
408        assert_eq!(session.phase(), SessionPhase::Initialized);
409
410        // Calling mark_initialized again should fail (already in target state)
411        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}