1use 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#[derive(Debug, Clone, PartialEq)]
20pub enum SessionState {
21 Idle,
23 Disconnected,
25 Connecting,
27 Connected,
29 Reconnecting,
31 Failed(String),
33}
34
35pub trait NotificationHandler: Send + Sync {
37 fn handle_notification(&self, notification: JsonRpcNotification);
39}
40
41#[derive(Debug, Clone)]
43pub struct SessionConfig {
44 pub auto_reconnect: bool,
46 pub max_reconnect_attempts: u32,
48 pub reconnect_delay_ms: u64,
50 pub max_reconnect_delay_ms: u64,
52 pub reconnect_backoff: f64,
54 pub connection_timeout_ms: u64,
56 pub heartbeat_interval_ms: u64,
58 pub heartbeat_timeout_ms: u64,
60
61 pub session_timeout: Duration,
64 pub request_timeout: Duration,
66 pub max_concurrent_requests: u32,
68 pub enable_compression: bool,
70 pub buffer_size: usize,
72
73 pub retry_config: RetryConfig,
76 pub enable_circuit_breaker: bool,
78 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(), enable_circuit_breaker: true,
100 circuit_breaker_config: CircuitBreakerConfig::default(),
101 }
102 }
103}
104
105pub struct ClientSession {
107 client: Arc<Mutex<McpClient>>,
109 config: SessionConfig,
111 state: Arc<RwLock<SessionState>>,
113 state_tx: watch::Sender<SessionState>,
115 state_rx: watch::Receiver<SessionState>,
117 notification_handlers: Arc<RwLock<Vec<Box<dyn NotificationHandler>>>>,
119 connected_at: Arc<RwLock<Option<Instant>>>,
121 reconnect_attempts: Arc<Mutex<u32>>,
123 shutdown_tx: Arc<Mutex<Option<mpsc::Sender<()>>>>,
125 #[allow(dead_code)]
127 retry_policy: Arc<RetryPolicy>,
128}
129
130impl ClientSession {
131 pub fn new(client: McpClient) -> Self {
133 let config = SessionConfig::default();
134 let (state_tx, state_rx) = watch::channel(SessionState::Disconnected);
135
136 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 pub fn with_config(client: McpClient, config: SessionConfig) -> Self {
162 let (state_tx, state_rx) = watch::channel(SessionState::Disconnected);
163
164 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 pub async fn state(&self) -> SessionState {
190 let state = self.state.read().await;
191 state.clone()
192 }
193
194 pub fn subscribe_state_changes(&self) -> watch::Receiver<SessionState> {
196 self.state_rx.clone()
197 }
198
199 pub async fn is_connected(&self) -> bool {
201 let state = self.state.read().await;
202 matches!(*state, SessionState::Connected)
203 }
204
205 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 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 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 {
244 let mut connected_at = self.connected_at.write().await;
245 *connected_at = Some(Instant::now());
246 }
247
248 {
250 let mut attempts = self.reconnect_attempts.lock().await;
251 *attempts = 0;
252 }
253
254 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 pub async fn disconnect(&self) -> McpResult<()> {
275 self.stop_background_tasks().await;
277
278 {
280 let client = self.client.lock().await;
281 client.disconnect().await?;
282 }
283
284 self.transition_state(SessionState::Disconnected).await?;
286
287 {
289 let mut connected_at = self.connected_at.write().await;
290 *connected_at = None;
291 }
292
293 Ok(())
294 }
295
296 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 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 self.transition_state(SessionState::Connected).await?;
339
340 {
342 let mut connected_at_guard = self.connected_at.write().await;
343 *connected_at_guard = Some(Instant::now());
344 }
345
346 {
348 let mut attempts = self.reconnect_attempts.lock().await;
349 *attempts = 0;
350 }
351
352 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 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 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 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 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 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 pub fn client(&self) -> Arc<Mutex<McpClient>> {
420 self.client.clone()
421 }
422
423 pub fn config(&self) -> &SessionConfig {
425 &self.config
426 }
427
428 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 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 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 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 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 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 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); }
480
481 {
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 }
505 Err(_) => {
506 break;
508 }
509 }
510 }
511 }
512 }
513 });
514 }
515
516 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 {
534 let current_state = state.read().await;
535 if !matches!(*current_state, SessionState::Connected) {
536 break;
537 }
538 }
539
540 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 let _ = state_tx.send(SessionState::Disconnected);
549 break;
550 }
551 }
552 }
553 }
554 });
555 }
556
557 Ok(())
558 }
559
560 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; }
570 }
571
572 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 if self.state_tx.send(new_state).is_err() {
581 }
583
584 Ok(())
585 }
586}
587
588pub 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
601pub struct ResourceUpdateHandler {
603 callback: Box<dyn Fn(String) + Send + Sync>,
604}
605
606impl ResourceUpdateHandler {
607 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
630pub struct ToolListChangedHandler {
632 callback: Box<dyn Fn() + Send + Sync>,
633}
634
635impl ToolListChangedHandler {
636 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
655pub struct ProgressHandler {
657 callback: Box<dyn Fn(String, f32, Option<u32>) + Send + Sync>,
658}
659
660impl ProgressHandler {
661 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#[derive(Debug, Clone)]
690pub struct SessionStats {
691 pub state: SessionState,
693 pub uptime: Option<Duration>,
695 pub reconnect_attempts: u32,
697 pub connected_at: Option<Instant>,
699}
700
701impl ClientSession {
702 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 struct MockTransport;
732
733 #[async_trait]
734 impl Transport for MockTransport {
735 async fn send_request(&mut self, _request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
736 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 let transport = MockTransport;
799 session.connect(transport).await.unwrap();
800 assert!(session.is_connected().await);
801
802 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 session
816 .add_notification_handler(LoggingNotificationHandler)
817 .await;
818
819 session
821 .add_notification_handler(ResourceUpdateHandler::new(|uri| {
822 println!("Resource updated: {uri}");
823 }))
824 .await;
825
826 session
828 .add_notification_handler(ToolListChangedHandler::new(|| {
829 println!("Tool list changed");
830 }))
831 .await;
832
833 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 assert_eq!(*state_rx.borrow(), SessionState::Disconnected);
881
882 session
884 .transition_state(SessionState::Connecting)
885 .await
886 .unwrap();
887
888 state_rx.changed().await.unwrap();
890 assert_eq!(*state_rx.borrow(), SessionState::Connecting);
891 }
892}