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}