1use crate::core::error::{McpError, McpResult};
9use crate::protocol::types::*;
10use serde::{Deserialize, Serialize};
11
12#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
18#[serde(transparent)]
19pub struct BatchRequest {
20 pub requests: Vec<BatchRequestItem>,
22}
23
24#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
26#[serde(untagged)]
27pub enum BatchRequestItem {
28 Request(JsonRpcRequest),
30 Notification(JsonRpcNotification),
32}
33
34#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
36#[serde(transparent)]
37pub struct BatchResponse {
38 pub responses: Vec<BatchResponseItem>,
40}
41
42#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
44#[serde(untagged)]
45pub enum BatchResponseItem {
46 Response(JsonRpcResponse),
48 Error(JsonRpcError),
50}
51
52impl BatchRequest {
57 pub fn new() -> Self {
59 Self {
60 requests: Vec::new(),
61 }
62 }
63
64 pub fn with_capacity(capacity: usize) -> Self {
66 Self {
67 requests: Vec::with_capacity(capacity),
68 }
69 }
70
71 pub fn add_request(mut self, request: JsonRpcRequest) -> Self {
73 self.requests.push(BatchRequestItem::Request(request));
74 self
75 }
76
77 pub fn add_notification(mut self, notification: JsonRpcNotification) -> Self {
79 self.requests
80 .push(BatchRequestItem::Notification(notification));
81 self
82 }
83
84 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 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 pub fn len(&self) -> usize {
110 self.requests.len()
111 }
112
113 pub fn is_empty(&self) -> bool {
115 self.requests.is_empty()
116 }
117
118 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 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 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 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
192impl BatchResponse {
197 pub fn new() -> Self {
199 Self {
200 responses: Vec::new(),
201 }
202 }
203
204 pub fn with_capacity(capacity: usize) -> Self {
206 Self {
207 responses: Vec::with_capacity(capacity),
208 }
209 }
210
211 pub fn add_response(mut self, response: JsonRpcResponse) -> Self {
213 self.responses.push(BatchResponseItem::Response(response));
214 self
215 }
216
217 pub fn add_error(mut self, error: JsonRpcError) -> Self {
219 self.responses.push(BatchResponseItem::Error(error));
220 self
221 }
222
223 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 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 pub fn len(&self) -> usize {
245 self.responses.len()
246 }
247
248 pub fn is_empty(&self) -> bool {
250 self.responses.is_empty()
251 }
252
253 pub fn validate(&self) -> McpResult<()> {
255 Ok(())
258 }
259
260 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 pub fn all_successful(&self) -> bool {
277 self.responses
278 .iter()
279 .all(|item| matches!(item, BatchResponseItem::Response(_)))
280 }
281
282 pub fn has_errors(&self) -> bool {
284 self.responses
285 .iter()
286 .any(|item| matches!(item, BatchResponseItem::Error(_)))
287 }
288
289 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
329pub struct BatchProcessor;
335
336impl BatchProcessor {
337 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 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 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 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#[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 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 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 let empty_batch = BatchRequest::new();
482 assert!(empty_batch.validate().is_err());
483
484 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 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}