Skip to main content

prism_mcp_rs/protocol/
batch.rs

1//! JSON-RPC Batch Request/Response Support (2025-11-25)
2//!
3//! Module provides support for JSON-RPC batch operations as defined in the
4//! JSON-RPC 2.0 specification, even though the MCP spec notes "simplified JSON-RPC
5//! without batching". Implementation is provided for completeness and future
6//! compatibility
7
8use crate::core::error::{McpError, McpResult};
9use crate::protocol::types::*;
10use serde::{Deserialize, Serialize};
11
12// ============================================================================
13// Batch Types
14// ============================================================================
15
16/// A JSON-RPC batch request containing multiple requests/notifications
17#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
18#[serde(transparent)]
19pub struct BatchRequest {
20    /// The individual requests in the batch
21    pub requests: Vec<BatchRequestItem>,
22}
23
24/// Individual item in a batch request
25#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
26#[serde(untagged)]
27pub enum BatchRequestItem {
28    /// A regular request expecting a response
29    Request(JsonRpcRequest),
30    /// A notification (no response expected)
31    Notification(JsonRpcNotification),
32}
33
34/// A JSON-RPC batch response containing multiple responses/errors
35#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
36#[serde(transparent)]
37pub struct BatchResponse {
38    /// The individual responses in the batch
39    pub responses: Vec<BatchResponseItem>,
40}
41
42/// Individual item in a batch response
43#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
44#[serde(untagged)]
45pub enum BatchResponseItem {
46    /// A successful response
47    Response(JsonRpcResponse),
48    /// An error response
49    Error(JsonRpcError),
50}
51
52// ============================================================================
53// Batch Request Implementation
54// ============================================================================
55
56impl BatchRequest {
57    /// Create a new empty batch request
58    pub fn new() -> Self {
59        Self {
60            requests: Vec::new(),
61        }
62    }
63
64    /// Create a batch request with initial capacity
65    pub fn with_capacity(capacity: usize) -> Self {
66        Self {
67            requests: Vec::with_capacity(capacity),
68        }
69    }
70
71    /// Add a request to the batch
72    pub fn add_request(mut self, request: JsonRpcRequest) -> Self {
73        self.requests.push(BatchRequestItem::Request(request));
74        self
75    }
76
77    /// Add a notification to the batch
78    pub fn add_notification(mut self, notification: JsonRpcNotification) -> Self {
79        self.requests
80            .push(BatchRequestItem::Notification(notification));
81        self
82    }
83
84    /// Add a method call as a request
85    pub fn add_call<T: Serialize>(
86        mut self,
87        id: RequestId,
88        method: String,
89        params: Option<T>,
90    ) -> McpResult<Self> {
91        let request = JsonRpcRequest::new(id, method, params)?;
92        self.requests.push(BatchRequestItem::Request(request));
93        Ok(self)
94    }
95
96    /// Add a method call as a notification
97    pub fn add_notify<T: Serialize>(
98        mut self,
99        method: String,
100        params: Option<T>,
101    ) -> McpResult<Self> {
102        let notification = JsonRpcNotification::new(method, params)?;
103        self.requests
104            .push(BatchRequestItem::Notification(notification));
105        Ok(self)
106    }
107
108    /// Get the number of requests in the batch
109    pub fn len(&self) -> usize {
110        self.requests.len()
111    }
112
113    /// Check if the batch is empty
114    pub fn is_empty(&self) -> bool {
115        self.requests.is_empty()
116    }
117
118    /// Validate the batch request
119    pub fn validate(&self) -> McpResult<()> {
120        if self.is_empty() {
121            return Err(McpError::Protocol(
122                "Batch request must contain at least one request".to_string(),
123            ));
124        }
125
126        // Check for duplicate IDs among requests
127        let mut seen_ids = std::collections::HashSet::new();
128        for item in &self.requests {
129            if let BatchRequestItem::Request(req) = item {
130                if let serde_json::Value::Null = req.id {
131                    // Null IDs are allowed
132                    continue;
133                }
134                let id_str = serde_json::to_string(&req.id)
135                    .map_err(|e| McpError::Protocol(format!("Invalid request ID: {e}")))?;
136                if !seen_ids.insert(id_str) {
137                    return Err(McpError::Protocol(format!(
138                        "Duplicate request ID in batch: {:?}",
139                        req.id
140                    )));
141                }
142            }
143        }
144
145        Ok(())
146    }
147
148    /// Split batch into requests and notifications
149    pub fn split(self) -> (Vec<JsonRpcRequest>, Vec<JsonRpcNotification>) {
150        let mut requests = Vec::new();
151        let mut notifications = Vec::new();
152
153        for item in self.requests {
154            match item {
155                BatchRequestItem::Request(req) => requests.push(req),
156                BatchRequestItem::Notification(notif) => notifications.push(notif),
157            }
158        }
159
160        (requests, notifications)
161    }
162}
163
164impl Default for BatchRequest {
165    fn default() -> Self {
166        Self::new()
167    }
168}
169
170impl From<Vec<JsonRpcRequest>> for BatchRequest {
171    fn from(requests: Vec<JsonRpcRequest>) -> Self {
172        Self {
173            requests: requests
174                .into_iter()
175                .map(BatchRequestItem::Request)
176                .collect(),
177        }
178    }
179}
180
181impl From<Vec<JsonRpcNotification>> for BatchRequest {
182    fn from(notifications: Vec<JsonRpcNotification>) -> Self {
183        Self {
184            requests: notifications
185                .into_iter()
186                .map(BatchRequestItem::Notification)
187                .collect(),
188        }
189    }
190}
191
192// ============================================================================
193// Batch Response Implementation
194// ============================================================================
195
196impl BatchResponse {
197    /// Create a new empty batch response
198    pub fn new() -> Self {
199        Self {
200            responses: Vec::new(),
201        }
202    }
203
204    /// Create a batch response with initial capacity
205    pub fn with_capacity(capacity: usize) -> Self {
206        Self {
207            responses: Vec::with_capacity(capacity),
208        }
209    }
210
211    /// Add a successful response to the batch
212    pub fn add_response(mut self, response: JsonRpcResponse) -> Self {
213        self.responses.push(BatchResponseItem::Response(response));
214        self
215    }
216
217    /// Add an error response to the batch
218    pub fn add_error(mut self, error: JsonRpcError) -> Self {
219        self.responses.push(BatchResponseItem::Error(error));
220        self
221    }
222
223    /// Add a success result
224    pub fn add_success<T: Serialize>(mut self, id: RequestId, result: T) -> McpResult<Self> {
225        let response = JsonRpcResponse::success(id, result)?;
226        self.responses.push(BatchResponseItem::Response(response));
227        Ok(self)
228    }
229
230    /// Add an error result
231    pub fn add_failure(
232        mut self,
233        id: RequestId,
234        code: i32,
235        message: String,
236        data: Option<serde_json::Value>,
237    ) -> Self {
238        let error = JsonRpcError::error(id, code, message, data);
239        self.responses.push(BatchResponseItem::Error(error));
240        self
241    }
242
243    /// Get the number of responses in the batch
244    pub fn len(&self) -> usize {
245        self.responses.len()
246    }
247
248    /// Check if the batch is empty
249    pub fn is_empty(&self) -> bool {
250        self.responses.is_empty()
251    }
252
253    /// Validate the batch response
254    pub fn validate(&self) -> McpResult<()> {
255        // Batch responses can be empty if all requests were notifications
256        // No duplicate ID check needed for responses
257        Ok(())
258    }
259
260    /// Split batch into successes and errors
261    pub fn split(self) -> (Vec<JsonRpcResponse>, Vec<JsonRpcError>) {
262        let mut successes = Vec::new();
263        let mut errors = Vec::new();
264
265        for item in self.responses {
266            match item {
267                BatchResponseItem::Response(resp) => successes.push(resp),
268                BatchResponseItem::Error(err) => errors.push(err),
269            }
270        }
271
272        (successes, errors)
273    }
274
275    /// Check if all responses are successful
276    pub fn all_successful(&self) -> bool {
277        self.responses
278            .iter()
279            .all(|item| matches!(item, BatchResponseItem::Response(_)))
280    }
281
282    /// Check if any response is an error
283    pub fn has_errors(&self) -> bool {
284        self.responses
285            .iter()
286            .any(|item| matches!(item, BatchResponseItem::Error(_)))
287    }
288
289    /// Get all error responses
290    pub fn errors(&self) -> Vec<&JsonRpcError> {
291        self.responses
292            .iter()
293            .filter_map(|item| {
294                if let BatchResponseItem::Error(err) = item {
295                    Some(err)
296                } else {
297                    None
298                }
299            })
300            .collect()
301    }
302}
303
304impl Default for BatchResponse {
305    fn default() -> Self {
306        Self::new()
307    }
308}
309
310impl From<Vec<JsonRpcResponse>> for BatchResponse {
311    fn from(responses: Vec<JsonRpcResponse>) -> Self {
312        Self {
313            responses: responses
314                .into_iter()
315                .map(BatchResponseItem::Response)
316                .collect(),
317        }
318    }
319}
320
321impl From<Vec<JsonRpcError>> for BatchResponse {
322    fn from(errors: Vec<JsonRpcError>) -> Self {
323        Self {
324            responses: errors.into_iter().map(BatchResponseItem::Error).collect(),
325        }
326    }
327}
328
329// ============================================================================
330// Batch Processor
331// ============================================================================
332
333/// Helper for processing batch requests
334pub struct BatchProcessor;
335
336impl BatchProcessor {
337    /// Process a batch request and return a batch response
338    pub async fn process<F, Fut>(batch: BatchRequest, handler: F) -> McpResult<BatchResponse>
339    where
340        F: Fn(JsonRpcRequest) -> Fut,
341        Fut: std::future::Future<Output = McpResult<serde_json::Value>>,
342    {
343        batch.validate()?;
344
345        let mut response = BatchResponse::new();
346        let (requests, _notifications) = batch.split();
347
348        // Process each request
349        // Note: Notifications don't get responses
350        for request in requests {
351            let id = request.id.clone();
352            match handler(request).await {
353                Ok(result) => {
354                    response = response.add_success(id, result)?;
355                }
356                Err(err) => {
357                    let (code, message) = match err {
358                        McpError::Protocol(msg) => (error_codes::INVALID_REQUEST, msg),
359                        McpError::MethodNotFound(msg) => (error_codes::METHOD_NOT_FOUND, msg),
360                        McpError::InvalidParams(msg) => (error_codes::INVALID_PARAMS, msg),
361                        _ => (error_codes::INTERNAL_ERROR, err.to_string()),
362                    };
363                    response = response.add_failure(id, code, message, None);
364                }
365            }
366        }
367
368        Ok(response)
369    }
370
371    /// Create a parse error response for invalid batch JSON
372    pub fn parse_error() -> BatchResponse {
373        BatchResponse {
374            responses: vec![BatchResponseItem::Error(JsonRpcError::error(
375                serde_json::Value::Null,
376                error_codes::PARSE_ERROR,
377                "Invalid JSON-RPC batch request".to_string(),
378                None,
379            ))],
380        }
381    }
382
383    /// Create an empty batch error response
384    pub fn empty_batch_error() -> BatchResponse {
385        BatchResponse {
386            responses: vec![BatchResponseItem::Error(JsonRpcError::error(
387                serde_json::Value::Null,
388                error_codes::INVALID_REQUEST,
389                "Batch request must not be empty".to_string(),
390                None,
391            ))],
392        }
393    }
394}
395
396// ============================================================================
397// Tests
398// ============================================================================
399
400#[cfg(test)]
401mod tests {
402    use super::*;
403    use serde_json::json;
404
405    #[test]
406    fn test_batch_request_creation() {
407        let batch = BatchRequest::new()
408            .add_call(
409                json!(1),
410                "method1".to_string(),
411                Some(json!({"key": "value"})),
412            )
413            .unwrap()
414            .add_notify("notification1".to_string(), Some(json!({"data": "test"})))
415            .unwrap();
416
417        assert_eq!(batch.len(), 2);
418        assert!(!batch.is_empty());
419
420        let (requests, notifications) = batch.split();
421        assert_eq!(requests.len(), 1);
422        assert_eq!(notifications.len(), 1);
423    }
424
425    #[test]
426    fn test_batch_response_creation() {
427        let batch = BatchResponse::new()
428            .add_success(json!(1), json!({"result": "success"}))
429            .unwrap()
430            .add_failure(
431                json!(2),
432                error_codes::METHOD_NOT_FOUND,
433                "Method not found".to_string(),
434                None,
435            );
436
437        assert_eq!(batch.len(), 2);
438        assert!(!batch.all_successful());
439        assert!(batch.has_errors());
440        assert_eq!(batch.errors().len(), 1);
441    }
442
443    #[test]
444    fn test_batch_serialization() {
445        // Test batch request serialization
446        let req1 = JsonRpcRequest {
447            jsonrpc: JSONRPC_VERSION.to_string(),
448            id: json!(1),
449            method: "test.method1".to_string(),
450            params: Some(json!({"param": "value"})),
451        };
452
453        let notif1 = JsonRpcNotification {
454            jsonrpc: JSONRPC_VERSION.to_string(),
455            method: "test.notify".to_string(),
456            params: Some(json!({"event": "occurred"})),
457        };
458
459        let batch = BatchRequest {
460            requests: vec![
461                BatchRequestItem::Request(req1),
462                BatchRequestItem::Notification(notif1),
463            ],
464        };
465
466        let json = serde_json::to_value(&batch).unwrap();
467        assert!(json.is_array());
468        assert_eq!(json[0]["id"], 1);
469        assert_eq!(json[0]["method"], "test.method1");
470        assert!(json[1]["id"].is_null());
471        assert_eq!(json[1]["method"], "test.notify");
472
473        // Test deserialization
474        let batch2: BatchRequest = serde_json::from_value(json).unwrap();
475        assert_eq!(batch.len(), batch2.len());
476    }
477
478    #[test]
479    fn test_batch_validation() {
480        // Empty batch should fail validation
481        let empty_batch = BatchRequest::new();
482        assert!(empty_batch.validate().is_err());
483
484        // Batch with duplicate IDs should fail
485        let duplicate_batch = BatchRequest::new()
486            .add_call(json!(1), "method1".to_string(), None::<()>)
487            .unwrap()
488            .add_call(json!(1), "method2".to_string(), None::<()>)
489            .unwrap();
490        assert!(duplicate_batch.validate().is_err());
491
492        // Valid batch should pass
493        let valid_batch = BatchRequest::new()
494            .add_call(json!(1), "method1".to_string(), None::<()>)
495            .unwrap()
496            .add_call(json!(2), "method2".to_string(), None::<()>)
497            .unwrap()
498            .add_notify("notify".to_string(), None::<()>)
499            .unwrap();
500        assert!(valid_batch.validate().is_ok());
501    }
502
503    #[test]
504    fn test_batch_response_helpers() {
505        let batch = BatchResponse::new()
506            .add_success(json!(1), json!({"data": "result1"}))
507            .unwrap()
508            .add_success(json!(2), json!({"data": "result2"}))
509            .unwrap();
510
511        assert!(batch.all_successful());
512        assert!(!batch.has_errors());
513        assert_eq!(batch.errors().len(), 0);
514
515        let batch_with_error = batch.add_failure(
516            json!(3),
517            error_codes::INTERNAL_ERROR,
518            "Internal error".to_string(),
519            None,
520        );
521
522        assert!(!batch_with_error.all_successful());
523        assert!(batch_with_error.has_errors());
524        assert_eq!(batch_with_error.errors().len(), 1);
525    }
526
527    #[tokio::test]
528    async fn test_batch_processor() {
529        let batch = BatchRequest::new()
530            .add_call(json!(1), "echo".to_string(), Some(json!({"msg": "hello"})))
531            .unwrap()
532            .add_call(json!(2), "echo".to_string(), Some(json!({"msg": "world"})))
533            .unwrap();
534
535        let response = BatchProcessor::process(batch, |req| async move {
536            if req.method == "echo" {
537                Ok(req.params.unwrap_or(json!({})))
538            } else {
539                Err(McpError::MethodNotFound(format!(
540                    "Unknown method: {}",
541                    req.method
542                )))
543            }
544        })
545        .await
546        .unwrap();
547
548        assert_eq!(response.len(), 2);
549        assert!(response.all_successful());
550    }
551
552    #[test]
553    fn test_special_batch_errors() {
554        let parse_error = BatchProcessor::parse_error();
555        assert_eq!(parse_error.len(), 1);
556        assert!(parse_error.has_errors());
557
558        let empty_error = BatchProcessor::empty_batch_error();
559        assert_eq!(empty_error.len(), 1);
560        assert!(empty_error.has_errors());
561    }
562}