Skip to main content

prism_mcp_rs/transport/
traits.rs

1//! Transport layer traits and abstractions
2//!
3//! Module defines the core transport traits that enable MCP communication
4//! over different protocols like STDIO, HTTP, and WebSocket.
5
6use crate::core::error::McpResult;
7use crate::protocol::types::{JsonRpcNotification, JsonRpcRequest, JsonRpcResponse};
8use async_trait::async_trait;
9use serde_json::Value;
10use tokio::sync::{mpsc, oneshot};
11
12/// Client-side handle for a long-lived `subscriptions/listen` request.
13pub struct ClientSubscription {
14    id: Value,
15    notifications: mpsc::UnboundedReceiver<JsonRpcNotification>,
16    completion: Option<oneshot::Receiver<McpResult<JsonRpcResponse>>>,
17    abort_handle: Option<tokio::task::AbortHandle>,
18}
19
20impl ClientSubscription {
21    pub fn new(
22        id: Value,
23        notifications: mpsc::UnboundedReceiver<JsonRpcNotification>,
24        completion: oneshot::Receiver<McpResult<JsonRpcResponse>>,
25    ) -> Self {
26        Self {
27            id,
28            notifications,
29            completion: Some(completion),
30            abort_handle: None,
31        }
32    }
33
34    #[cfg(feature = "http")]
35    pub(crate) fn with_abort_handle(mut self, handle: tokio::task::AbortHandle) -> Self {
36        self.abort_handle = Some(handle);
37        self
38    }
39
40    pub fn id(&self) -> &Value {
41        &self.id
42    }
43
44    /// Wait for the next notification on this subscription.
45    pub async fn next(&mut self) -> Option<JsonRpcNotification> {
46        self.notifications.recv().await
47    }
48
49    /// Wait for a graceful terminal response. HTTP connection closure may
50    /// instead end this channel without a response.
51    pub async fn completion(&mut self) -> McpResult<JsonRpcResponse> {
52        let receiver = self.completion.take().ok_or_else(|| {
53            crate::core::error::McpError::Protocol(
54                "subscription completion was already consumed".to_string(),
55            )
56        })?;
57        receiver.await.map_err(|_| {
58            crate::core::error::McpError::Transport(
59                "subscription stream closed without a final response".to_string(),
60            )
61        })?
62    }
63}
64
65impl Drop for ClientSubscription {
66    fn drop(&mut self) {
67        if let Some(handle) = self.abort_handle.take() {
68            handle.abort();
69        }
70    }
71}
72
73/// Transport trait for MCP clients
74///
75/// **Note**: This trait is primarily for internal use and advanced custom transport implementations.
76/// Most users should use the provided transport implementations (StdioTransport, HttpTransport, WebSocketTransport)
77/// rather than implementing this trait directly.
78///
79/// ## When to use Transport directly:
80/// - Implementing a custom transport protocol (e.g., IPC, named pipes, custom network protocol)
81/// - Creating mock transports for testing
82/// - Building transport middleware or decorators
83///
84/// ## Example Custom Transport:
85/// ```no_run
86/// use prism_mcp_rs::transport::traits::Transport;
87/// use prism_mcp_rs::protocol::{JsonRpcRequest, JsonRpcResponse, JsonRpcNotification};
88/// use prism_mcp_rs::core::error::McpResult;
89/// use async_trait::async_trait;
90///
91/// struct CustomTransport {
92///     // Your transport implementation
93/// }
94///
95/// #[async_trait]
96/// impl Transport for CustomTransport {
97///     async fn send_request(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
98///         // Implementation
99///         # todo!()
100///     }
101///     
102///     async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()> {
103///         // Implementation
104///         # todo!()
105///     }
106///     
107///     async fn receive_notification(&mut self) -> McpResult<Option<JsonRpcNotification>> {
108///         // Implementation
109///         # todo!()
110///     }
111///     
112///     async fn close(&mut self) -> McpResult<()> {
113///         // Close implementation
114///         Ok(())
115///     }
116/// }
117/// ```
118///
119/// Trait defines the interface for sending requests and receiving responses
120/// in a client-side MCP connection.
121#[async_trait]
122pub trait Transport: Send + Sync {
123    /// Send a JSON-RPC request and wait for a response
124    ///
125    /// # Arguments
126    /// * `request` - The JSON-RPC request to send
127    ///
128    /// # Returns
129    /// Result containing the JSON-RPC response or an error
130    async fn send_request(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse>;
131
132    /// Send a JSON-RPC notification (no response expected)
133    ///
134    /// # Arguments
135    /// * `notification` - The JSON-RPC notification to send
136    ///
137    /// # Returns
138    /// Result indicating success or an error
139    async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()>;
140
141    /// Receive a notification from the server (non-blocking)
142    ///
143    /// # Returns
144    /// Result containing an optional notification or an error
145    async fn receive_notification(&mut self) -> McpResult<Option<JsonRpcNotification>>;
146
147    /// Open a modern long-lived subscription stream.
148    async fn open_subscription(
149        &mut self,
150        _request: JsonRpcRequest,
151    ) -> McpResult<ClientSubscription> {
152        Err(crate::core::error::McpError::MethodNotFound(
153            "subscriptions/listen is not supported by this transport".to_string(),
154        ))
155    }
156
157    /// Cancel an open subscription. HTTP implementations close the response
158    /// stream; STDIO implementations send `notifications/cancelled`.
159    async fn cancel_subscription(&mut self, _request_id: &Value) -> McpResult<()> {
160        Err(crate::core::error::McpError::MethodNotFound(
161            "subscription cancellation is not supported by this transport".to_string(),
162        ))
163    }
164
165    /// Handle incoming server request (for bidirectional communication)
166    ///
167    /// Method is called to check for incoming requests from the server.
168    /// It should be non-blocking and return None if no request is available.
169    ///
170    /// # Returns
171    /// Result containing an optional server request or an error
172    async fn handle_incoming_request(&mut self) -> McpResult<Option<JsonRpcRequest>> {
173        // Default implementation - no bidirectional support
174        Ok(None)
175    }
176
177    /// Send response to server request (for bidirectional communication)
178    ///
179    /// Method sends a response back to the server for a server-initiated request.
180    ///
181    /// # Arguments
182    /// * `response` - The JSON-RPC response to send
183    ///
184    /// # Returns
185    /// Result indicating success or an error
186    async fn send_response(&mut self, _response: JsonRpcResponse) -> McpResult<()> {
187        // Default implementation - no bidirectional support
188        Err(crate::core::error::McpError::MethodNotFound(
189            "Bidirectional communication not supported by this transport".to_string(),
190        ))
191    }
192
193    /// Close the transport connection
194    ///
195    /// # Returns
196    /// Result indicating success or an error
197    async fn close(&mut self) -> McpResult<()>;
198
199    /// Check if the transport is connected
200    ///
201    /// # Returns
202    /// True if the transport is connected and ready for communication
203    fn is_connected(&self) -> bool {
204        true // Default implementation - assume connected
205    }
206
207    /// Get connection information for debugging
208    ///
209    /// # Returns
210    /// String describing the connection
211    fn connection_info(&self) -> String {
212        "Unknown transport".to_string()
213    }
214}
215
216/// Server request handler function type
217pub type ServerRequestHandler = std::sync::Arc<
218    dyn Fn(
219            JsonRpcRequest,
220        ) -> std::pin::Pin<
221            Box<dyn std::future::Future<Output = McpResult<JsonRpcResponse>> + Send + 'static>,
222        > + Send
223        + Sync,
224>;
225
226/// Transport trait for MCP servers
227///
228/// Trait defines the interface for handling incoming requests and
229/// sending responses in a server-side MCP connection.
230#[async_trait]
231pub trait ServerTransport: Send + Sync {
232    /// Provide registered tool input schemas to transports that mirror and
233    /// validate MCP routing headers. Other transports may ignore them.
234    fn set_tool_schemas(
235        &mut self,
236        _schemas: std::collections::HashMap<String, serde_json::Value>,
237    ) -> McpResult<()> {
238        Ok(())
239    }
240
241    /// Provide the capabilities used to acknowledge subscription filters.
242    fn set_server_capabilities(
243        &mut self,
244        _capabilities: crate::protocol::ServerCapabilities,
245    ) -> McpResult<()> {
246        Ok(())
247    }
248
249    /// Attach task-status notifications to the transport's subscription hub.
250    fn set_task_notifications(
251        &mut self,
252        _receiver: tokio::sync::broadcast::Receiver<JsonRpcNotification>,
253    ) -> McpResult<()> {
254        Ok(())
255    }
256
257    /// Start the server transport and begin listening for connections
258    ///
259    /// # Returns
260    /// Result indicating success or an error
261    async fn start(&mut self) -> McpResult<()>;
262
263    /// Set the request handler that will process incoming requests
264    ///
265    /// # Arguments
266    /// * `handler` - The request handler function
267    fn set_request_handler(&mut self, handler: ServerRequestHandler);
268
269    /// Send a JSON-RPC notification to the client
270    ///
271    /// # Arguments
272    /// * `notification` - The JSON-RPC notification to send
273    ///
274    /// # Returns
275    /// Result indicating success or an error
276    async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()>;
277
278    /// Stop the server transport
279    ///
280    /// # Returns
281    /// Result indicating success or an error
282    async fn stop(&mut self) -> McpResult<()>;
283
284    /// Check if the server is running
285    ///
286    /// # Returns
287    /// True if the server is running and accepting connections
288    fn is_running(&self) -> bool {
289        true // Default implementation - assume running
290    }
291
292    /// Get server information for debugging
293    ///
294    /// # Returns
295    /// String describing the server state
296    fn server_info(&self) -> String {
297        "Unknown server transport".to_string()
298    }
299}
300
301/// Transport configuration options
302#[derive(Debug, Clone)]
303pub struct TransportConfig {
304    /// Connection timeout in milliseconds
305    pub connect_timeout_ms: Option<u64>,
306    /// Read timeout in milliseconds
307    pub read_timeout_ms: Option<u64>,
308    /// Write timeout in milliseconds
309    pub write_timeout_ms: Option<u64>,
310    /// Maximum message size in bytes
311    pub max_message_size: Option<usize>,
312    /// Keep-alive interval in milliseconds
313    pub keep_alive_ms: Option<u64>,
314    /// Whether to enable compression
315    pub compression: bool,
316    /// Custom headers for HTTP-based transports
317    pub headers: std::collections::HashMap<String, String>,
318}
319
320impl Default for TransportConfig {
321    fn default() -> Self {
322        Self {
323            connect_timeout_ms: Some(30_000),         // 30 seconds
324            read_timeout_ms: Some(60_000),            // 60 seconds
325            write_timeout_ms: Some(30_000),           // 30 seconds
326            max_message_size: Some(16 * 1024 * 1024), // 16 MB
327            keep_alive_ms: Some(30_000),              // 30 seconds
328            compression: false,
329            headers: std::collections::HashMap::new(),
330        }
331    }
332}
333
334/// Connection state for transports
335#[derive(Debug, Clone, PartialEq)]
336pub enum ConnectionState {
337    /// Transport is disconnected
338    Disconnected,
339    /// Transport is connecting
340    Connecting,
341    /// Transport is connected and ready
342    Connected,
343    /// Transport is reconnecting after an error
344    Reconnecting,
345    /// Transport is closing
346    Closing,
347    /// Transport has encountered an error
348    Error(String),
349}
350
351/// Transport statistics for monitoring
352#[derive(Debug, Clone, Default)]
353pub struct TransportStats {
354    /// Number of requests sent
355    pub requests_sent: u64,
356    /// Number of responses received
357    pub responses_received: u64,
358    /// Number of notifications sent
359    pub notifications_sent: u64,
360    /// Number of notifications received
361    pub notifications_received: u64,
362    /// Number of connection errors
363    pub connection_errors: u64,
364    /// Number of protocol errors
365    pub protocol_errors: u64,
366    /// Total bytes sent
367    pub bytes_sent: u64,
368    /// Total bytes received
369    pub bytes_received: u64,
370    /// Connection uptime in milliseconds
371    pub uptime_ms: u64,
372}
373
374/// Trait for transports that support statistics
375pub trait TransportStats_: Send + Sync {
376    /// Get current transport statistics
377    fn stats(&self) -> TransportStats;
378
379    /// Reset transport statistics
380    fn reset_stats(&mut self);
381}
382
383/// Trait for transports that support reconnection
384#[async_trait]
385pub trait ReconnectableTransport: Transport {
386    /// Attempt to reconnect the transport
387    ///
388    /// # Returns
389    /// Result indicating success or an error
390    async fn reconnect(&mut self) -> McpResult<()>;
391
392    /// Set the reconnection configuration
393    ///
394    /// # Arguments
395    /// * `config` - Reconnection configuration
396    fn set_reconnect_config(&mut self, config: ReconnectConfig);
397
398    /// Get the current connection state
399    fn connection_state(&self) -> ConnectionState;
400}
401
402/// Configuration for automatic reconnection
403#[derive(Debug, Clone)]
404pub struct ReconnectConfig {
405    /// Whether automatic reconnection is enabled
406    pub enabled: bool,
407    /// Maximum number of reconnection attempts
408    pub max_attempts: Option<u32>,
409    /// Initial delay before first reconnection attempt (milliseconds)
410    pub initial_delay_ms: u64,
411    /// Maximum delay between reconnection attempts (milliseconds)
412    pub max_delay_ms: u64,
413    /// Multiplier for exponential backoff
414    pub backoff_multiplier: f64,
415    /// Jitter factor for randomizing delays (0.0 to 1.0)
416    pub jitter_factor: f64,
417}
418
419impl Default for ReconnectConfig {
420    fn default() -> Self {
421        Self {
422            enabled: true,
423            max_attempts: Some(5),
424            initial_delay_ms: 1000, // 1 second
425            max_delay_ms: 30_000,   // 30 seconds
426            backoff_multiplier: 2.0,
427            jitter_factor: 0.1,
428        }
429    }
430}
431
432/// Trait for transports that support message filtering
433pub trait FilterableTransport: Send + Sync {
434    /// Set a message filter function
435    ///
436    /// # Arguments
437    /// * `filter` - Function that returns true if message should be processed
438    fn set_message_filter(&mut self, filter: Box<dyn Fn(&JsonRpcRequest) -> bool + Send + Sync>);
439
440    /// Clear the message filter
441    fn clear_message_filter(&mut self);
442}
443
444/// Transport event for monitoring and debugging
445#[derive(Debug, Clone)]
446pub enum TransportEvent {
447    /// Connection established
448    Connected,
449    /// Connection lost
450    Disconnected,
451    /// Message sent
452    MessageSent {
453        /// Message type
454        message_type: String,
455        /// Message size in bytes
456        size: usize,
457    },
458    /// Message received
459    MessageReceived {
460        /// Message type
461        message_type: String,
462        /// Message size in bytes
463        size: usize,
464    },
465    /// Error occurred
466    Error {
467        /// Error message
468        message: String,
469    },
470}
471
472/// Trait for transports that support event listeners
473pub trait EventEmittingTransport: Send + Sync {
474    /// Add an event listener
475    ///
476    /// # Arguments
477    /// * `listener` - Event listener function
478    fn add_event_listener(&mut self, listener: Box<dyn Fn(TransportEvent) + Send + Sync>);
479
480    /// Remove all event listeners
481    fn clear_event_listeners(&mut self);
482}
483
484#[cfg(test)]
485mod tests {
486    use super::*;
487
488    #[test]
489    fn test_transport_config_default() {
490        let config = TransportConfig::default();
491        assert_eq!(config.connect_timeout_ms, Some(30_000));
492        assert_eq!(config.read_timeout_ms, Some(60_000));
493        assert_eq!(config.max_message_size, Some(16 * 1024 * 1024));
494        assert!(!config.compression);
495    }
496
497    #[test]
498    fn test_reconnect_config_default() {
499        let config = ReconnectConfig::default();
500        assert!(config.enabled);
501        assert_eq!(config.max_attempts, Some(5));
502        assert_eq!(config.initial_delay_ms, 1000);
503        assert_eq!(config.max_delay_ms, 30_000);
504        assert_eq!(config.backoff_multiplier, 2.0);
505        assert_eq!(config.jitter_factor, 0.1);
506    }
507
508    #[test]
509    fn test_connection_state_equality() {
510        assert_eq!(ConnectionState::Connected, ConnectionState::Connected);
511        assert_eq!(ConnectionState::Disconnected, ConnectionState::Disconnected);
512        assert_ne!(ConnectionState::Connected, ConnectionState::Disconnected);
513
514        let error1 = ConnectionState::Error("test".to_string());
515        let error2 = ConnectionState::Error("test".to_string());
516        let error3 = ConnectionState::Error("other".to_string());
517        assert_eq!(error1, error2);
518        assert_ne!(error1, error3);
519    }
520
521    #[test]
522    fn test_transport_stats_default() {
523        let stats = TransportStats::default();
524        assert_eq!(stats.requests_sent, 0);
525        assert_eq!(stats.responses_received, 0);
526        assert_eq!(stats.bytes_sent, 0);
527        assert_eq!(stats.bytes_received, 0);
528    }
529}