Skip to main content

prism_mcp_rs/client/
session.rs

1//! Client session management
2//!
3//! Module provides session management for MCP clients, including connection
4//! state tracking, notification handling, and automatic reconnection capabilities.
5
6use std::sync::Arc;
7use std::time::{Duration, Instant};
8use tokio::sync::{broadcast, mpsc, watch, Mutex, RwLock};
9use tokio::time::timeout;
10
11use crate::client::mcp_client::McpClient;
12use crate::core::error::{McpError, McpResult};
13use crate::core::logging::ErrorContext;
14use crate::core::retry::{CircuitBreakerConfig, RetryConfig, RetryPolicy};
15use crate::protocol::{messages::*, methods, types::*, ConnectResult};
16use crate::transport::traits::Transport;
17
18/// Session state
19#[derive(Debug, Clone, PartialEq)]
20pub enum SessionState {
21    /// Session is idle (ready but not actively processing)
22    Idle,
23    /// Session is disconnected
24    Disconnected,
25    /// Session is connecting
26    Connecting,
27    /// Session is connected and active
28    Connected,
29    /// Session is reconnecting after a failure
30    Reconnecting,
31    /// Session has failed and cannot reconnect
32    Failed(String),
33}
34
35/// Notification handler trait
36pub trait NotificationHandler: Send + Sync {
37    /// Handle a notification from the server
38    fn handle_notification(&self, notification: JsonRpcNotification);
39}
40
41/// Session configuration
42#[derive(Debug, Clone)]
43pub struct SessionConfig {
44    /// Whether to enable automatic reconnection
45    pub auto_reconnect: bool,
46    /// Maximum number of reconnection attempts
47    pub max_reconnect_attempts: u32,
48    /// Initial reconnection delay in milliseconds
49    pub reconnect_delay_ms: u64,
50    /// Maximum reconnection delay in milliseconds
51    pub max_reconnect_delay_ms: u64,
52    /// Reconnection backoff multiplier
53    pub reconnect_backoff: f64,
54    /// Connection timeout in milliseconds
55    pub connection_timeout_ms: u64,
56    /// Heartbeat interval in milliseconds (0 to disable)
57    pub heartbeat_interval_ms: u64,
58    /// Heartbeat timeout in milliseconds
59    pub heartbeat_timeout_ms: u64,
60
61    // Additional fields for test compatibility
62    /// Session timeout duration
63    pub session_timeout: Duration,
64    /// Request timeout duration
65    pub request_timeout: Duration,
66    /// Maximum concurrent requests
67    pub max_concurrent_requests: u32,
68    /// Enable compression
69    pub enable_compression: bool,
70    /// Buffer size for operations
71    pub buffer_size: usize,
72
73    // Production-ready retry configuration
74    /// Retry policy for operations
75    pub retry_config: RetryConfig,
76    /// Enable circuit breaker for resilience
77    pub enable_circuit_breaker: bool,
78    /// Circuit breaker configuration
79    pub circuit_breaker_config: CircuitBreakerConfig,
80}
81
82impl Default for SessionConfig {
83    fn default() -> Self {
84        Self {
85            auto_reconnect: true,
86            max_reconnect_attempts: 5,
87            reconnect_delay_ms: 1000,
88            max_reconnect_delay_ms: 30000,
89            reconnect_backoff: 2.0,
90            connection_timeout_ms: 10000,
91            heartbeat_interval_ms: 30000,
92            heartbeat_timeout_ms: 5000,
93            session_timeout: Duration::from_secs(300),
94            request_timeout: Duration::from_secs(30),
95            max_concurrent_requests: 10,
96            enable_compression: false,
97            buffer_size: 8192,
98            retry_config: RetryConfig::network(), // Use network-improved retry config
99            enable_circuit_breaker: true,
100            circuit_breaker_config: CircuitBreakerConfig::default(),
101        }
102    }
103}
104
105/// Client session that manages connection lifecycle and notifications
106pub struct ClientSession {
107    /// The underlying MCP client
108    client: Arc<Mutex<McpClient>>,
109    /// Session configuration
110    config: SessionConfig,
111    /// Current session state
112    state: Arc<RwLock<SessionState>>,
113    /// State change broadcaster
114    state_tx: watch::Sender<SessionState>,
115    /// State change receiver
116    state_rx: watch::Receiver<SessionState>,
117    /// Notification handlers
118    notification_handlers: Arc<RwLock<Vec<Box<dyn NotificationHandler>>>>,
119    /// Connection timestamp
120    connected_at: Arc<RwLock<Option<Instant>>>,
121    /// Reconnection attempts counter
122    reconnect_attempts: Arc<Mutex<u32>>,
123    /// Shutdown signal
124    shutdown_tx: Arc<Mutex<Option<mpsc::Sender<()>>>>,
125    /// Retry policy for connection operations
126    #[allow(dead_code)]
127    retry_policy: Arc<RetryPolicy>,
128}
129
130impl ClientSession {
131    /// Create a new client session
132    pub fn new(client: McpClient) -> Self {
133        let config = SessionConfig::default();
134        let (state_tx, state_rx) = watch::channel(SessionState::Disconnected);
135
136        // Create retry policy based on configuration
137        let retry_policy = if config.enable_circuit_breaker {
138            Arc::new(RetryPolicy::with_circuit_breaker(
139                config.retry_config.clone(),
140                config.circuit_breaker_config.clone(),
141            ))
142        } else {
143            Arc::new(RetryPolicy::new(config.retry_config.clone()))
144        };
145
146        Self {
147            client: Arc::new(Mutex::new(client)),
148            config,
149            state: Arc::new(RwLock::new(SessionState::Disconnected)),
150            state_tx,
151            state_rx,
152            notification_handlers: Arc::new(RwLock::new(Vec::new())),
153            connected_at: Arc::new(RwLock::new(None)),
154            reconnect_attempts: Arc::new(Mutex::new(0)),
155            shutdown_tx: Arc::new(Mutex::new(None)),
156            retry_policy,
157        }
158    }
159
160    /// Create a new client session with custom configuration
161    pub fn with_config(client: McpClient, config: SessionConfig) -> Self {
162        let (state_tx, state_rx) = watch::channel(SessionState::Disconnected);
163
164        // Create retry policy based on configuration
165        let retry_policy = if config.enable_circuit_breaker {
166            Arc::new(RetryPolicy::with_circuit_breaker(
167                config.retry_config.clone(),
168                config.circuit_breaker_config.clone(),
169            ))
170        } else {
171            Arc::new(RetryPolicy::new(config.retry_config.clone()))
172        };
173
174        Self {
175            client: Arc::new(Mutex::new(client)),
176            config,
177            state: Arc::new(RwLock::new(SessionState::Disconnected)),
178            state_tx,
179            state_rx,
180            notification_handlers: Arc::new(RwLock::new(Vec::new())),
181            connected_at: Arc::new(RwLock::new(None)),
182            reconnect_attempts: Arc::new(Mutex::new(0)),
183            shutdown_tx: Arc::new(Mutex::new(None)),
184            retry_policy,
185        }
186    }
187
188    /// Get the current session state
189    pub async fn state(&self) -> SessionState {
190        let state = self.state.read().await;
191        state.clone()
192    }
193
194    /// Subscribe to state changes
195    pub fn subscribe_state_changes(&self) -> watch::Receiver<SessionState> {
196        self.state_rx.clone()
197    }
198
199    /// Check if the session is connected
200    pub async fn is_connected(&self) -> bool {
201        let state = self.state.read().await;
202        matches!(*state, SessionState::Connected)
203    }
204
205    /// Get connection uptime
206    pub async fn uptime(&self) -> Option<Duration> {
207        let connected_at = self.connected_at.read().await;
208        connected_at.map(|time| time.elapsed())
209    }
210
211    /// Add a notification handler
212    pub async fn add_notification_handler<H>(&self, handler: H)
213    where
214        H: NotificationHandler + 'static,
215    {
216        let mut handlers = self.notification_handlers.write().await;
217        handlers.push(Box::new(handler));
218    }
219
220    /// Connect to the server with the provided transport
221    pub async fn connect<T>(&self, transport: T) -> McpResult<ConnectResult>
222    where
223        T: Transport + 'static,
224    {
225        self.transition_state(SessionState::Connecting).await?;
226
227        let connect_future = async {
228            let mut client = self.client.lock().await;
229            client.connect(transport).await
230        };
231
232        let result = timeout(
233            Duration::from_millis(self.config.connection_timeout_ms),
234            connect_future,
235        )
236        .await;
237
238        match result {
239            Ok(Ok(init_result)) => {
240                self.transition_state(SessionState::Connected).await?;
241
242                // Record connection time
243                {
244                    let mut connected_at = self.connected_at.write().await;
245                    *connected_at = Some(Instant::now());
246                }
247
248                // Reset reconnection attempts
249                {
250                    let mut attempts = self.reconnect_attempts.lock().await;
251                    *attempts = 0;
252                }
253
254                // Start background tasks
255                self.start_background_tasks().await?;
256
257                Ok(init_result)
258            }
259            Ok(Err(error)) => {
260                self.transition_state(SessionState::Failed(error.to_string()))
261                    .await?;
262                Err(error)
263            }
264            Err(_) => {
265                let error = McpError::Connection("Connection timeout".to_string());
266                self.transition_state(SessionState::Failed(error.to_string()))
267                    .await?;
268                Err(error)
269            }
270        }
271    }
272
273    /// Disconnect from the server
274    pub async fn disconnect(&self) -> McpResult<()> {
275        // Stop background tasks
276        self.stop_background_tasks().await;
277
278        // Disconnect the client
279        {
280            let client = self.client.lock().await;
281            client.disconnect().await?;
282        }
283
284        // Update state
285        self.transition_state(SessionState::Disconnected).await?;
286
287        // Clear connection time
288        {
289            let mut connected_at = self.connected_at.write().await;
290            *connected_at = None;
291        }
292
293        Ok(())
294    }
295
296    /// Reconnect to the server with smart retry logic
297    pub async fn reconnect<T>(
298        &self,
299        transport_factory: impl Fn() -> T + Send + Sync + 'static,
300    ) -> McpResult<ConnectResult>
301    where
302        T: Transport + 'static,
303    {
304        if !self.config.auto_reconnect {
305            let error = McpError::connection("Auto-reconnect is disabled");
306            return Err(error);
307        }
308
309        let _context = ErrorContext::new("session_reconnect")
310            .with_component("client_session")
311            .with_extra(
312                "max_attempts",
313                serde_json::Value::from(self.config.max_reconnect_attempts),
314            );
315
316        self.transition_state(SessionState::Reconnecting).await?;
317
318        // Simple retry loop with smart error handling
319        let mut last_error = None;
320
321        for attempt in 1..=self.config.retry_config.max_attempts {
322            let transport = transport_factory();
323
324            let connect_future = async {
325                let mut client_guard = self.client.lock().await;
326                client_guard.connect(transport).await
327            };
328
329            let result = tokio::time::timeout(
330                Duration::from_millis(self.config.connection_timeout_ms),
331                connect_future,
332            )
333            .await;
334
335            match result {
336                Ok(Ok(init_result)) => {
337                    // Success! Update state and reset counters
338                    self.transition_state(SessionState::Connected).await?;
339
340                    // Record connection time
341                    {
342                        let mut connected_at_guard = self.connected_at.write().await;
343                        *connected_at_guard = Some(Instant::now());
344                    }
345
346                    // Reset reconnection attempts
347                    {
348                        let mut attempts = self.reconnect_attempts.lock().await;
349                        *attempts = 0;
350                    }
351
352                    // Start background tasks for the new connection
353                    if let Err(e) = self.start_background_tasks().await {
354                        tracing::warn!("Failed to start background tasks after reconnect: {}", e);
355                    }
356
357                    return Ok(init_result);
358                }
359                Ok(Err(error)) => {
360                    last_error = Some(error.clone());
361
362                    // Check if error is recoverable and if we should retry
363                    if !error.is_recoverable() || attempt >= self.config.retry_config.max_attempts {
364                        self.transition_state(SessionState::Failed(error.to_string()))
365                            .await?;
366                        return Err(error);
367                    }
368
369                    // Log retry attempt
370                    tracing::warn!(
371                        "Reconnection attempt {} failed: {} (recoverable: {}, will retry: {})",
372                        attempt,
373                        error,
374                        error.is_recoverable(),
375                        attempt < self.config.retry_config.max_attempts
376                    );
377                }
378                Err(_) => {
379                    // Timeout occurred
380                    let timeout_error = McpError::timeout("Connection timeout during reconnect");
381                    last_error = Some(timeout_error.clone());
382
383                    if attempt >= self.config.retry_config.max_attempts {
384                        self.transition_state(SessionState::Failed(timeout_error.to_string()))
385                            .await?;
386                        return Err(timeout_error);
387                    }
388
389                    tracing::warn!("Reconnection attempt {} timed out (will retry)", attempt);
390                }
391            }
392
393            // Apply delay before next attempt (except for last attempt)
394            if attempt < self.config.retry_config.max_attempts {
395                let delay_ms = self.config.retry_config.initial_delay_ms
396                    * (self
397                        .config
398                        .retry_config
399                        .backoff_multiplier
400                        .powi(attempt as i32 - 1) as u64);
401                let delay =
402                    Duration::from_millis(delay_ms.min(self.config.retry_config.max_delay_ms));
403
404                tracing::debug!("Waiting {:?} before next reconnection attempt", delay);
405                tokio::time::sleep(delay).await;
406            }
407        }
408
409        // All retries exhausted
410        let final_error = last_error
411            .unwrap_or_else(|| McpError::internal("Reconnection failed without capturing error"));
412
413        self.transition_state(SessionState::Failed(final_error.to_string()))
414            .await?;
415        Err(final_error)
416    }
417
418    /// Get the underlying client (for direct operations)
419    pub fn client(&self) -> Arc<Mutex<McpClient>> {
420        self.client.clone()
421    }
422
423    /// Get session configuration
424    pub fn config(&self) -> &SessionConfig {
425        &self.config
426    }
427
428    // ========================================================================
429    // Convenience Methods for Common Operations
430    // ========================================================================
431
432    /// List available tools from the server
433    pub async fn list_tools(&self, cursor: Option<String>) -> McpResult<ListToolsResult> {
434        let client = self.client.lock().await;
435        client.list_tools(cursor).await
436    }
437
438    /// Call a tool on the server
439    pub async fn call_tool(&self, params: CallToolParams) -> McpResult<CallToolResult> {
440        let client = self.client.lock().await;
441        client.call_tool(params.name, params.arguments).await
442    }
443
444    /// List available resources from the server
445    pub async fn list_resources(&self, cursor: Option<String>) -> McpResult<ListResourcesResult> {
446        let client = self.client.lock().await;
447        client.list_resources(cursor).await
448    }
449
450    /// Read a resource from the server
451    pub async fn read_resource(&self, params: ReadResourceParams) -> McpResult<ReadResourceResult> {
452        let client = self.client.lock().await;
453        client.read_resource(params.uri).await
454    }
455
456    /// List available prompts from the server
457    pub async fn list_prompts(&self, cursor: Option<String>) -> McpResult<ListPromptsResult> {
458        let client = self.client.lock().await;
459        client.list_prompts(cursor).await
460    }
461
462    /// Get a prompt from the server
463    pub async fn get_prompt(&self, params: GetPromptParams) -> McpResult<GetPromptResult> {
464        let client = self.client.lock().await;
465        client.get_prompt(params.name, params.arguments).await
466    }
467
468    // ========================================================================
469    // Background Tasks
470    // ========================================================================
471
472    /// Start background tasks (notification handling, heartbeat)
473    async fn start_background_tasks(&self) -> McpResult<()> {
474        let (_shutdown_tx, shutdown_rx): (broadcast::Sender<()>, broadcast::Receiver<()>) =
475            broadcast::channel(16);
476        {
477            let mut shutdown_guard = self.shutdown_tx.lock().await;
478            *shutdown_guard = Some(mpsc::channel(1).0); // Store a dummy for interface compatibility
479        }
480
481        // Start notification handler task
482        {
483            let client = self.client.clone();
484            let handlers = self.notification_handlers.clone();
485            let mut shutdown_rx_clone = shutdown_rx.resubscribe();
486
487            tokio::spawn(async move {
488                loop {
489                    tokio::select! {
490                        _ = shutdown_rx_clone.recv() => break,
491                        notification_result = async {
492                            let client_guard = client.lock().await;
493                            client_guard.receive_notification().await
494                        } => {
495                            match notification_result {
496                                Ok(Some(notification)) => {
497                                    let handlers_guard = handlers.read().await;
498                                    for handler in handlers_guard.iter() {
499                                        handler.handle_notification(notification.clone());
500                                    }
501                                }
502                                Ok(None) => {
503                                    // No notification available, continue
504                                }
505                                Err(_) => {
506                                    // Error receiving notification, might be disconnected
507                                    break;
508                                }
509                            }
510                        }
511                    }
512                }
513            });
514        }
515
516        // Start heartbeat task if enabled
517        if self.config.heartbeat_interval_ms > 0 {
518            let client = self.client.clone();
519            let heartbeat_interval = Duration::from_millis(self.config.heartbeat_interval_ms);
520            let heartbeat_timeout = Duration::from_millis(self.config.heartbeat_timeout_ms);
521            let state = self.state.clone();
522            let state_tx = self.state_tx.clone();
523            let mut shutdown_rx_clone = shutdown_rx.resubscribe();
524
525            tokio::spawn(async move {
526                let mut interval = tokio::time::interval(heartbeat_interval);
527
528                loop {
529                    tokio::select! {
530                        _ = shutdown_rx_clone.recv() => break,
531                        _ = interval.tick() => {
532                            // Check if we're still connected
533                            {
534                                let current_state = state.read().await;
535                                if !matches!(*current_state, SessionState::Connected) {
536                                    break;
537                                }
538                            }
539
540                            // Send ping
541                            let ping_result = timeout(heartbeat_timeout, async {
542                                let client_guard = client.lock().await;
543                                client_guard.ping().await
544                            }).await;
545
546                            if ping_result.is_err() {
547                                // Heartbeat failed, mark as disconnected
548                                let _ = state_tx.send(SessionState::Disconnected);
549                                break;
550                            }
551                        }
552                    }
553                }
554            });
555        }
556
557        Ok(())
558    }
559
560    /// Stop background tasks
561    async fn stop_background_tasks(&self) {
562        let shutdown_tx = {
563            let mut shutdown_guard = self.shutdown_tx.lock().await;
564            shutdown_guard.take()
565        };
566
567        if let Some(tx) = shutdown_tx {
568            let _ = tx.send(()).await; // Ignore error if receiver is dropped
569        }
570    }
571
572    /// Transition to a new state
573    async fn transition_state(&self, new_state: SessionState) -> McpResult<()> {
574        {
575            let mut state = self.state.write().await;
576            *state = new_state.clone();
577        }
578
579        // Broadcast the state change
580        if self.state_tx.send(new_state).is_err() {
581            // Receiver may have been dropped, which is okay
582        }
583
584        Ok(())
585    }
586}
587
588/// Default notification handler that logs notifications
589pub struct LoggingNotificationHandler;
590
591impl NotificationHandler for LoggingNotificationHandler {
592    fn handle_notification(&self, notification: JsonRpcNotification) {
593        tracing::info!(
594            "Received notification: {} {:?}",
595            notification.method,
596            notification.params
597        );
598    }
599}
600
601/// Resource update notification handler
602pub struct ResourceUpdateHandler {
603    callback: Box<dyn Fn(String) + Send + Sync>,
604}
605
606impl ResourceUpdateHandler {
607    /// Create a new resource update handler
608    pub fn new<F>(callback: F) -> Self
609    where
610        F: Fn(String) + Send + Sync + 'static,
611    {
612        Self {
613            callback: Box::new(callback),
614        }
615    }
616}
617
618impl NotificationHandler for ResourceUpdateHandler {
619    fn handle_notification(&self, notification: JsonRpcNotification) {
620        if notification.method == methods::RESOURCES_UPDATED {
621            if let Some(params) = notification.params {
622                if let Ok(update_params) = serde_json::from_value::<ResourceUpdatedParams>(params) {
623                    (self.callback)(update_params.uri);
624                }
625            }
626        }
627    }
628}
629
630/// Tool list changed notification handler
631pub struct ToolListChangedHandler {
632    callback: Box<dyn Fn() + Send + Sync>,
633}
634
635impl ToolListChangedHandler {
636    /// Create a new tool list changed handler
637    pub fn new<F>(callback: F) -> Self
638    where
639        F: Fn() + Send + Sync + 'static,
640    {
641        Self {
642            callback: Box::new(callback),
643        }
644    }
645}
646
647impl NotificationHandler for ToolListChangedHandler {
648    fn handle_notification(&self, notification: JsonRpcNotification) {
649        if notification.method == methods::TOOLS_LIST_CHANGED {
650            (self.callback)();
651        }
652    }
653}
654
655/// Progress notification handler
656pub struct ProgressHandler {
657    callback: Box<dyn Fn(String, f32, Option<u32>) + Send + Sync>,
658}
659
660impl ProgressHandler {
661    /// Create a new progress handler
662    pub fn new<F>(callback: F) -> Self
663    where
664        F: Fn(String, f32, Option<u32>) + Send + Sync + 'static,
665    {
666        Self {
667            callback: Box::new(callback),
668        }
669    }
670}
671
672impl NotificationHandler for ProgressHandler {
673    fn handle_notification(&self, notification: JsonRpcNotification) {
674        if notification.method == methods::PROGRESS {
675            if let Some(params) = notification.params {
676                if let Ok(progress_params) = serde_json::from_value::<ProgressParams>(params) {
677                    (self.callback)(
678                        progress_params.progress_token.to_string(),
679                        progress_params.progress,
680                        progress_params.total.map(|t| t as u32),
681                    );
682                }
683            }
684        }
685    }
686}
687
688/// Session statistics
689#[derive(Debug, Clone)]
690pub struct SessionStats {
691    /// Current session state
692    pub state: SessionState,
693    /// Connection uptime
694    pub uptime: Option<Duration>,
695    /// Number of reconnection attempts
696    pub reconnect_attempts: u32,
697    /// Connection timestamp
698    pub connected_at: Option<Instant>,
699}
700
701impl ClientSession {
702    /// Get session statistics
703    pub async fn stats(&self) -> SessionStats {
704        let state = self.state().await;
705        let uptime = self.uptime().await;
706        let reconnect_attempts = {
707            let attempts = self.reconnect_attempts.lock().await;
708            *attempts
709        };
710        let connected_at = {
711            let connected_at = self.connected_at.read().await;
712            *connected_at
713        };
714
715        SessionStats {
716            state,
717            uptime,
718            reconnect_attempts,
719            connected_at,
720        }
721    }
722}
723
724#[cfg(test)]
725mod tests {
726    use super::*;
727    use crate::client::mcp_client::McpClient;
728    use async_trait::async_trait;
729
730    // Mock transport for testing
731    struct MockTransport;
732
733    #[async_trait]
734    impl Transport for MockTransport {
735        async fn send_request(&mut self, _request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
736            // Return a successful initialize response
737            let init_result = InitializeResult::new(
738                crate::protocol::LATEST_PROTOCOL_VERSION.to_string(),
739                ServerCapabilities::default(),
740                ServerInfo {
741                    name: "test-server".to_string(),
742                    version: "1.0.0".to_string(),
743                    description: None,
744                    title: Some("Test Server".to_string()),
745                    website_url: None,
746                    icons: None,
747                },
748            );
749            JsonRpcResponse::success(serde_json::Value::from(1), init_result)
750                .map_err(|e| McpError::Serialization(e.to_string()))
751        }
752
753        async fn send_notification(&mut self, _notification: JsonRpcNotification) -> McpResult<()> {
754            Ok(())
755        }
756
757        async fn receive_notification(&mut self) -> McpResult<Option<JsonRpcNotification>> {
758            Ok(None)
759        }
760
761        async fn close(&mut self) -> McpResult<()> {
762            Ok(())
763        }
764    }
765
766    #[tokio::test]
767    async fn test_session_creation() {
768        let client = McpClient::new("test-client".to_string(), "1.0.0".to_string());
769        let session = ClientSession::new(client);
770
771        assert_eq!(session.state().await, SessionState::Disconnected);
772        assert!(!session.is_connected().await);
773        assert!(session.uptime().await.is_none());
774    }
775
776    #[tokio::test]
777    async fn test_session_connection() {
778        let mut client = McpClient::new("test-client".to_string(), "1.0.0".to_string());
779        client.set_protocol_mode(crate::protocol::ProtocolMode::LegacyOnly);
780        let session = ClientSession::new(client);
781
782        let transport = MockTransport;
783        let result = session.connect(transport).await;
784
785        assert!(result.is_ok());
786        assert_eq!(session.state().await, SessionState::Connected);
787        assert!(session.is_connected().await);
788        assert!(session.uptime().await.is_some());
789    }
790
791    #[tokio::test]
792    async fn test_session_disconnect() {
793        let mut client = McpClient::new("test-client".to_string(), "1.0.0".to_string());
794        client.set_protocol_mode(crate::protocol::ProtocolMode::LegacyOnly);
795        let session = ClientSession::new(client);
796
797        // Connect first
798        let transport = MockTransport;
799        session.connect(transport).await.unwrap();
800        assert!(session.is_connected().await);
801
802        // Then disconnect
803        session.disconnect().await.unwrap();
804        assert_eq!(session.state().await, SessionState::Disconnected);
805        assert!(!session.is_connected().await);
806        assert!(session.uptime().await.is_none());
807    }
808
809    #[tokio::test]
810    async fn test_notification_handlers() {
811        let client = McpClient::new("test-client".to_string(), "1.0.0".to_string());
812        let session = ClientSession::new(client);
813
814        // Add a logging notification handler
815        session
816            .add_notification_handler(LoggingNotificationHandler)
817            .await;
818
819        // Add a resource update handler
820        session
821            .add_notification_handler(ResourceUpdateHandler::new(|uri| {
822                println!("Resource updated: {uri}");
823            }))
824            .await;
825
826        // Add a tool list changed handler
827        session
828            .add_notification_handler(ToolListChangedHandler::new(|| {
829                println!("Tool list changed");
830            }))
831            .await;
832
833        // Add a progress handler
834        session
835            .add_notification_handler(ProgressHandler::new(|token, progress, total| {
836                println!("Progress {token}: {progress} / {total:?}");
837            }))
838            .await;
839
840        let handlers = session.notification_handlers.read().await;
841        assert_eq!(handlers.len(), 4);
842    }
843
844    #[tokio::test]
845    async fn test_session_stats() {
846        let client = McpClient::new("test-client".to_string(), "1.0.0".to_string());
847        let session = ClientSession::new(client);
848
849        let stats = session.stats().await;
850        assert_eq!(stats.state, SessionState::Disconnected);
851        assert!(stats.uptime.is_none());
852        assert_eq!(stats.reconnect_attempts, 0);
853        assert!(stats.connected_at.is_none());
854    }
855
856    #[tokio::test]
857    async fn test_session_config() {
858        let client = McpClient::new("test-client".to_string(), "1.0.0".to_string());
859        let config = SessionConfig {
860            auto_reconnect: false,
861            max_reconnect_attempts: 10,
862            reconnect_delay_ms: 2000,
863            ..Default::default()
864        };
865        let session = ClientSession::with_config(client, config.clone());
866
867        assert!(!session.config().auto_reconnect);
868        assert_eq!(session.config().max_reconnect_attempts, 10);
869        assert_eq!(session.config().reconnect_delay_ms, 2000);
870    }
871
872    #[tokio::test]
873    async fn test_state_subscription() {
874        let client = McpClient::new("test-client".to_string(), "1.0.0".to_string());
875        let session = ClientSession::new(client);
876
877        let mut state_rx = session.subscribe_state_changes();
878
879        // Initial state
880        assert_eq!(*state_rx.borrow(), SessionState::Disconnected);
881
882        // Change state
883        session
884            .transition_state(SessionState::Connecting)
885            .await
886            .unwrap();
887
888        // Wait for change
889        state_rx.changed().await.unwrap();
890        assert_eq!(*state_rx.borrow(), SessionState::Connecting);
891    }
892}