Skip to main content

prism_mcp_rs/server/
async_methods.rs

1//! Async request handling methods for MCP server
2//!
3//! This module provides additional async methods for handling requests,
4//! processing messages, and managing server operations asynchronously.
5
6use crate::core::error::{McpError, McpResult};
7use crate::protocol::{
8    JsonRpcError, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse,
9};
10use crate::server::McpServer;
11use futures::future::Future;
12use serde_json::Value;
13use std::pin::Pin;
14
15impl McpServer {
16    /// Process any JSON-RPC message asynchronously
17    ///
18    /// This method handles all types of JSON-RPC messages (requests, responses,
19    /// notifications, errors) and returns an appropriate response.
20    ///
21    /// # Arguments
22    /// * `msg` - The JSON-RPC message to process
23    ///
24    /// # Returns
25    /// A result containing the response message or an error
26    ///
27    /// # Examples
28    /// ```no_run
29    /// # use prism_mcp_rs::server::McpServer;
30    /// # use prism_mcp_rs::protocol::{JsonRpcMessage, JsonRpcRequest};
31    /// # use serde_json::json;
32    /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
33    /// let server = McpServer::new("server".to_string(), "1.0.0".to_string());
34    /// let request = JsonRpcRequest::new(
35    ///     json!(1),
36    ///     "ping".to_string(),
37    ///     None::<serde_json::Value>
38    /// )?;
39    /// let message = JsonRpcMessage::Request(request);
40    /// let response = server.process_message(message).await?;
41    /// # Ok(())
42    /// # }
43    /// ```
44    pub async fn process_message(&self, msg: JsonRpcMessage) -> McpResult<JsonRpcMessage> {
45        match msg {
46            JsonRpcMessage::Request(req) => {
47                let response = self.handle_request(req).await?;
48                Ok(JsonRpcMessage::Response(response))
49            }
50            JsonRpcMessage::Notification(notif) => {
51                self.handle_notification(notif).await?;
52                // Notifications don't get responses
53                Ok(JsonRpcMessage::Response(
54                    JsonRpcResponse::success_unchecked(Value::Null, Value::Null),
55                ))
56            }
57            JsonRpcMessage::Response(resp) => {
58                // Responses are typically handled by the client side
59                // For server, we just echo them back
60                Ok(JsonRpcMessage::Response(resp))
61            }
62            JsonRpcMessage::Error(err) => {
63                // Error messages are also typically client-side
64                Ok(JsonRpcMessage::Error(err))
65            }
66        }
67    }
68
69    /// Handle a JSON-RPC notification asynchronously
70    ///
71    /// Notifications are one-way messages that don't expect a response.
72    ///
73    /// # Arguments
74    /// * `notification` - The notification to handle
75    ///
76    /// # Returns
77    /// A result indicating success or failure of handling the notification
78    pub async fn handle_notification(&self, notification: JsonRpcNotification) -> McpResult<()> {
79        match notification.method.as_str() {
80            "initialized" => {
81                // Client has acknowledged initialization
82                tracing::info!("Client initialized successfully");
83                Ok(())
84            }
85            "cancelled" => {
86                // Request cancellation notification
87                if let Some(params) = notification.params {
88                    tracing::info!("Request cancelled: {:?}", params);
89                }
90                Ok(())
91            }
92            "$/cancelRequest" => {
93                // LSP-style cancellation
94                if let Some(params) = notification.params {
95                    tracing::info!("LSP cancel request: {:?}", params);
96                }
97                Ok(())
98            }
99            "$/setTrace" => {
100                // Set trace level notification
101                if let Some(params) = notification.params {
102                    tracing::info!("Set trace level: {:?}", params);
103                }
104                Ok(())
105            }
106            _ => {
107                tracing::warn!("Unknown notification method: {}", notification.method);
108                Ok(())
109            }
110        }
111    }
112
113    /// Handle multiple requests in parallel
114    ///
115    /// This method processes multiple requests concurrently and returns
116    /// all responses once they're complete.
117    ///
118    /// # Arguments
119    /// * `requests` - A vector of JSON-RPC requests to process
120    ///
121    /// # Returns
122    /// A vector of results, one for each request
123    ///
124    /// # Examples
125    /// ```no_run
126    /// # use prism_mcp_rs::server::McpServer;
127    /// # use prism_mcp_rs::protocol::JsonRpcRequest;
128    /// # use serde_json::json;
129    /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
130    /// let server = McpServer::new("server".to_string(), "1.0.0".to_string());
131    /// let requests = vec![
132    ///     JsonRpcRequest::new(json!(1), "ping".to_string(), None::<serde_json::Value>)?,
133    ///     JsonRpcRequest::new(json!(2), "tools/list".to_string(), None::<serde_json::Value>)?,
134    /// ];
135    /// let responses = server.handle_requests_parallel(requests).await;
136    /// # Ok(())
137    /// # }
138    /// ```
139    pub async fn handle_requests_parallel(
140        &self,
141        requests: Vec<JsonRpcRequest>,
142    ) -> Vec<McpResult<JsonRpcResponse>> {
143        use futures::future::join_all;
144
145        let futures: Vec<_> = requests
146            .into_iter()
147            .map(|req| self.handle_request(req))
148            .collect();
149
150        join_all(futures).await
151    }
152
153    /// Handle requests with a custom async processor
154    ///
155    /// This method allows you to provide a custom processor function that
156    /// can modify or filter requests before they're handled.
157    ///
158    /// # Arguments
159    /// * `request` - The request to process
160    /// * `processor` - An async function that processes the request
161    ///
162    /// # Returns
163    /// The processed response
164    pub async fn handle_request_with_processor<F, Fut>(
165        &self,
166        request: JsonRpcRequest,
167        processor: F,
168    ) -> McpResult<JsonRpcResponse>
169    where
170        F: FnOnce(JsonRpcRequest) -> Fut,
171        Fut: Future<Output = McpResult<JsonRpcRequest>>,
172    {
173        let processed_request = processor(request).await?;
174        self.handle_request(processed_request).await
175    }
176
177    /// Run the server with a custom message processor
178    ///
179    /// This method allows you to run the server with a custom async message
180    /// processor that can handle incoming messages with custom logic.
181    ///
182    /// # Arguments
183    /// * `processor` - An async function that processes messages
184    ///
185    /// # Returns
186    /// A result indicating success or failure
187    pub async fn run_with_processor<F, Fut>(&self, _processor: F) -> McpResult<()>
188    where
189        F: FnMut(JsonRpcMessage) -> Fut + Send + 'static,
190        Fut: Future<Output = McpResult<JsonRpcMessage>> + Send,
191    {
192        // Check if server is running
193        if !self.is_running().await {
194            return Err(McpError::Transport("Server not running".to_string()));
195        }
196
197        // In a real implementation, this would:
198        // 1. Receive messages from transport
199        // 2. Process them with the custom processor
200        // 3. Send responses back through transport
201
202        tracing::info!("Server running with custom processor");
203        Ok(())
204    }
205
206    /// Handle a batch of JSON-RPC requests
207    ///
208    /// This method processes a batch of requests according to the JSON-RPC
209    /// batch specification.
210    ///
211    /// # Arguments
212    /// * `batch` - A JSON array of requests
213    ///
214    /// # Returns
215    /// A JSON array of responses
216    pub async fn handle_batch(&self, batch: Vec<Value>) -> McpResult<Vec<Value>> {
217        use futures::future::join_all;
218
219        let futures: Vec<_> = batch
220            .into_iter()
221            .map(|value| async move {
222                match serde_json::from_value::<JsonRpcRequest>(value) {
223                    Ok(request) => {
224                        match self.handle_request(request).await {
225                            Ok(response) => serde_json::to_value(response).unwrap_or(Value::Null),
226                            Err(err) => {
227                                // Create an error response for internal error
228                                let error_response = JsonRpcError {
229                                    jsonrpc: "2.0".to_string(),
230                                    id: Value::Null,
231                                    error: crate::protocol::types::ErrorObject {
232                                        code: -32603, // Internal error
233                                        message: err.to_string(),
234                                        data: None,
235                                    },
236                                };
237                                serde_json::to_value(error_response).unwrap_or(Value::Null)
238                            }
239                        }
240                    }
241                    Err(err) => {
242                        // Create an error response for parse error
243                        let error_response = JsonRpcError {
244                            jsonrpc: "2.0".to_string(),
245                            id: Value::Null,
246                            error: crate::protocol::types::ErrorObject {
247                                code: -32700, // Parse error
248                                message: err.to_string(),
249                                data: None,
250                            },
251                        };
252                        serde_json::to_value(error_response).unwrap_or(Value::Null)
253                    }
254                }
255            })
256            .collect();
257
258        Ok(join_all(futures).await)
259    }
260
261    /// Stream responses for long-running operations
262    ///
263    /// This method provides a way to stream responses for operations that
264    /// may take a long time to complete.
265    ///
266    /// # Arguments
267    /// * `request` - The request to process
268    /// * `progress_callback` - A callback function for progress updates
269    ///
270    /// # Returns
271    /// The final response
272    pub async fn handle_request_streaming<F>(
273        &self,
274        request: JsonRpcRequest,
275        mut progress_callback: F,
276    ) -> McpResult<JsonRpcResponse>
277    where
278        F: FnMut(f32, String) + Send,
279    {
280        // Start processing
281        progress_callback(0.0, "Starting request processing".to_string());
282
283        // For demonstration, we'll just handle normally
284        // In a real implementation, this would integrate with
285        // long-running operations and provide progress updates
286        progress_callback(50.0, "Processing request".to_string());
287
288        let response = self.handle_request(request).await?;
289
290        progress_callback(100.0, "Request completed".to_string());
291
292        Ok(response)
293    }
294
295    /// Handle a request with timeout
296    ///
297    /// This method processes a request with a specified timeout.
298    ///
299    /// # Arguments
300    /// * `request` - The request to process
301    /// * `timeout` - The timeout duration in milliseconds
302    ///
303    /// # Returns
304    /// The response or a timeout error
305    pub async fn handle_request_with_timeout(
306        &self,
307        request: JsonRpcRequest,
308        timeout_ms: u64,
309    ) -> McpResult<JsonRpcResponse> {
310        use tokio::time::{timeout, Duration};
311
312        match timeout(
313            Duration::from_millis(timeout_ms),
314            self.handle_request(request),
315        )
316        .await
317        {
318            Ok(result) => result,
319            Err(_) => Err(McpError::Timeout(format!(
320                "Request timed out after {}ms",
321                timeout_ms
322            ))),
323        }
324    }
325
326    /// Handle a request with retry logic
327    ///
328    /// This method attempts to process a request with automatic retry
329    /// on failure.
330    ///
331    /// # Arguments
332    /// * `request` - The request to process
333    /// * `max_retries` - Maximum number of retry attempts
334    /// * `retry_delay_ms` - Delay between retries in milliseconds
335    ///
336    /// # Returns
337    /// The response or the last error after all retries
338    pub async fn handle_request_with_retry(
339        &self,
340        request: JsonRpcRequest,
341        max_retries: u32,
342        retry_delay_ms: u64,
343    ) -> McpResult<JsonRpcResponse> {
344        use tokio::time::{sleep, Duration};
345
346        let mut last_error = None;
347
348        for attempt in 0..=max_retries {
349            match self.handle_request(request.clone()).await {
350                Ok(response) => return Ok(response),
351                Err(err) => {
352                    last_error = Some(err);
353                    if attempt < max_retries {
354                        tracing::warn!(
355                            "Request failed (attempt {}/{}), retrying in {}ms: {}",
356                            attempt + 1,
357                            max_retries + 1,
358                            retry_delay_ms,
359                            last_error.as_ref().unwrap()
360                        );
361                        sleep(Duration::from_millis(retry_delay_ms)).await;
362                    }
363                }
364            }
365        }
366
367        Err(last_error.unwrap_or_else(|| McpError::internal("Unknown error")))
368    }
369
370    /// Process a message with middleware chain
371    ///
372    /// This method allows you to process messages through a chain of
373    /// middleware functions.
374    ///
375    /// # Arguments
376    /// * `message` - The message to process
377    /// * `middleware` - A vector of middleware functions
378    ///
379    /// # Returns
380    /// The processed message
381    pub async fn process_with_middleware<M>(
382        &self,
383        message: JsonRpcMessage,
384        middleware: Vec<M>,
385    ) -> McpResult<JsonRpcMessage>
386    where
387        M: Fn(JsonRpcMessage) -> Pin<Box<dyn Future<Output = McpResult<JsonRpcMessage>> + Send>>
388            + Send
389            + Sync,
390    {
391        let mut current = message;
392
393        for mw in middleware.iter() {
394            current = mw(current).await?;
395        }
396
397        self.process_message(current).await
398    }
399}
400
401#[cfg(test)]
402mod tests {
403    use super::*;
404    use serde_json::json;
405
406    #[tokio::test]
407    async fn test_process_message() {
408        let server = McpServer::new("test".to_string(), "1.0.0".to_string());
409
410        let request =
411            JsonRpcRequest::new(json!(1), "ping".to_string(), None::<serde_json::Value>).unwrap();
412
413        let message = JsonRpcMessage::Request(request);
414        let response = server.process_message(message).await.unwrap();
415
416        match response {
417            JsonRpcMessage::Response(_) => {}
418            _ => panic!("Expected response message"),
419        }
420    }
421
422    #[tokio::test]
423    async fn test_handle_notification() {
424        let server = McpServer::new("test".to_string(), "1.0.0".to_string());
425
426        let notification = JsonRpcNotification {
427            jsonrpc: "2.0".to_string(),
428            method: "initialized".to_string(),
429            params: None,
430        };
431
432        let result = server.handle_notification(notification).await;
433        assert!(result.is_ok());
434    }
435
436    #[tokio::test]
437    async fn test_parallel_requests() {
438        let server = McpServer::new("test".to_string(), "1.0.0".to_string());
439
440        let requests = vec![
441            JsonRpcRequest::new(json!(1), "ping".to_string(), None::<serde_json::Value>).unwrap(),
442            JsonRpcRequest::new(json!(2), "ping".to_string(), None::<serde_json::Value>).unwrap(),
443        ];
444
445        let responses = server.handle_requests_parallel(requests).await;
446        assert_eq!(responses.len(), 2);
447        assert!(responses.iter().all(|r| r.is_ok()));
448    }
449
450    #[tokio::test]
451    async fn test_request_with_timeout() {
452        let server = McpServer::new("test".to_string(), "1.0.0".to_string());
453
454        let request =
455            JsonRpcRequest::new(json!(1), "ping".to_string(), None::<serde_json::Value>).unwrap();
456
457        // Should complete within timeout
458        let result = server.handle_request_with_timeout(request, 1000).await;
459        assert!(result.is_ok());
460    }
461}