1use std::convert::Infallible;
19use std::fmt;
20use std::future::Future;
21use std::pin::Pin;
22#[cfg(any(feature = "http", feature = "websocket"))]
23use std::sync::Arc;
24use std::task::{Context, Poll};
25
26use pin_project_lite::pin_project;
27
28use tower::util::BoxCloneService;
29use tower_service::Service;
30
31use crate::error::JsonRpcError;
32use crate::protocol::{McpRequest, RequestId};
33#[cfg(any(feature = "http", feature = "websocket"))]
34use crate::router::McpRouter;
35use crate::router::{RouterRequest, RouterResponse, ToolAnnotationsMap};
36
37pub type McpBoxService = BoxCloneService<RouterRequest, RouterResponse, Infallible>;
44
45#[cfg(any(feature = "http", feature = "websocket"))]
52pub(crate) type ServiceFactory = Arc<dyn Fn(McpRouter) -> McpBoxService + Send + Sync>;
53
54#[cfg(any(feature = "http", feature = "websocket"))]
59pub(crate) fn identity_factory() -> ServiceFactory {
60 Arc::new(|router: McpRouter| {
61 let annotations = router.tool_annotations_map();
62 BoxCloneService::new(InjectAnnotations::new(router, annotations))
63 })
64}
65
66#[derive(Clone)]
73pub struct InjectAnnotations<S> {
74 inner: S,
75 annotations: ToolAnnotationsMap,
76}
77
78impl<S> InjectAnnotations<S> {
79 pub fn new(inner: S, annotations: ToolAnnotationsMap) -> Self {
81 Self { inner, annotations }
82 }
83}
84
85impl<S: fmt::Debug> fmt::Debug for InjectAnnotations<S> {
86 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
87 f.debug_struct("InjectAnnotations")
88 .field("inner", &self.inner)
89 .finish()
90 }
91}
92
93impl<S> Service<RouterRequest> for InjectAnnotations<S>
94where
95 S: Service<RouterRequest, Response = RouterResponse>,
96{
97 type Response = RouterResponse;
98 type Error = S::Error;
99 type Future = S::Future;
100
101 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
102 self.inner.poll_ready(cx)
103 }
104
105 fn call(&mut self, mut req: RouterRequest) -> Self::Future {
106 if matches!(&req.inner, McpRequest::CallTool(_)) {
107 req.extensions.insert(self.annotations.clone());
108 }
109 self.inner.call(req)
110 }
111}
112
113pub struct CatchError<S> {
125 inner: S,
126}
127
128impl<S> CatchError<S> {
129 pub fn new(inner: S) -> Self {
131 Self { inner }
132 }
133}
134
135impl<S: Clone> Clone for CatchError<S> {
136 fn clone(&self) -> Self {
137 Self {
138 inner: self.inner.clone(),
139 }
140 }
141}
142
143impl<S: fmt::Debug> fmt::Debug for CatchError<S> {
144 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
145 f.debug_struct("CatchError")
146 .field("inner", &self.inner)
147 .finish()
148 }
149}
150
151pin_project! {
152 pub struct CatchErrorFuture<F> {
154 #[pin]
155 inner: F,
156 request_id: Option<RequestId>,
157 }
158}
159
160impl<F, E> Future for CatchErrorFuture<F>
161where
162 F: Future<Output = Result<RouterResponse, E>>,
163 E: fmt::Display,
164{
165 type Output = Result<RouterResponse, Infallible>;
166
167 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
168 let this = self.project();
169 match this.inner.poll(cx) {
170 Poll::Pending => Poll::Pending,
171 Poll::Ready(Ok(response)) => Poll::Ready(Ok(response)),
172 Poll::Ready(Err(err)) => {
173 let request_id = this.request_id.take().unwrap_or(RequestId::Number(0));
174 Poll::Ready(Ok(RouterResponse {
175 id: request_id,
176 inner: Err(JsonRpcError::internal_error(err.to_string())),
177 }))
178 }
179 }
180 }
181}
182
183impl<S> Service<RouterRequest> for CatchError<S>
184where
185 S: Service<RouterRequest, Response = RouterResponse> + Clone + Send + 'static,
186 S::Error: fmt::Display + Send,
187 S::Future: Send,
188{
189 type Response = RouterResponse;
190 type Error = Infallible;
191 type Future = CatchErrorFuture<tower::util::Oneshot<S, RouterRequest>>;
192
193 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
194 Poll::Ready(Ok(()))
198 }
199
200 fn call(&mut self, req: RouterRequest) -> Self::Future {
201 let request_id = req.id.clone();
204 let fut = tower::ServiceExt::oneshot(self.inner.clone(), req);
205
206 CatchErrorFuture {
207 inner: fut,
208 request_id: Some(request_id),
209 }
210 }
211}
212
213#[cfg(test)]
214mod tests {
215 use std::sync::Arc;
216 use std::sync::atomic::{AtomicBool, Ordering};
217
218 use super::*;
219 use crate::jsonrpc::JsonRpcService;
220 use crate::protocol::{
221 CallToolParams, CallToolResult, JsonRpcMessage, JsonRpcRequest, JsonRpcResponse,
222 JsonRpcResponseMessage, RequestId, ToolAnnotations,
223 };
224 use crate::router::McpRouter;
225
226 #[derive(Clone)]
227 struct RejectReadinessOnce {
228 inner: McpRouter,
229 reject_next: Arc<AtomicBool>,
230 }
231
232 impl Service<RouterRequest> for RejectReadinessOnce {
233 type Response = RouterResponse;
234 type Error = std::io::Error;
235 type Future =
236 Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send + 'static>>;
237
238 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
239 if self.reject_next.swap(false, Ordering::SeqCst) {
240 return Poll::Ready(Err(std::io::Error::other("middleware was not ready")));
241 }
242 Service::poll_ready(&mut self.inner, cx).map_err(|never| match never {})
243 }
244
245 fn call(&mut self, request: RouterRequest) -> Self::Future {
246 let future = Service::call(&mut self.inner, request);
247 Box::pin(async move { Ok(future.await.expect("MCP router service is infallible")) })
248 }
249 }
250
251 #[test]
252 #[cfg(any(feature = "http", feature = "websocket"))]
253 fn test_identity_factory_produces_service() {
254 let router = McpRouter::new().server_info("test", "1.0.0");
255 let factory = identity_factory();
256 let _service = factory(router);
257 }
258
259 #[tokio::test]
260 async fn test_catch_error_passes_through_success() {
261 let router = McpRouter::new().server_info("test", "1.0.0");
262 let mut service = CatchError::new(router);
263
264 let req = RouterRequest {
265 id: RequestId::Number(1),
266 inner: crate::protocol::McpRequest::Ping,
267 extensions: crate::router::Extensions::new(),
268 };
269
270 let result = Service::call(&mut service, req).await;
271 assert!(result.is_ok());
272 let response = result.unwrap();
273 assert!(response.inner.is_ok());
274 }
275
276 #[tokio::test]
277 async fn readiness_error_is_correlated_once_through_jsonrpc_service() {
278 let router = McpRouter::new().server_info("test", "1.0.0");
279 let reject_next = Arc::new(AtomicBool::new(true));
280 let mut service = JsonRpcService::new(CatchError::new(RejectReadinessOnce {
281 inner: router,
282 reject_next,
283 }));
284
285 std::future::poll_fn(|cx| Service::<JsonRpcRequest>::poll_ready(&mut service, cx))
286 .await
287 .expect("adapter is infallible");
288 let first = Service::<JsonRpcRequest>::call(&mut service, JsonRpcRequest::new(1, "ping"))
289 .await
290 .expect("JSON-RPC service call");
291 let JsonRpcResponse::Error(first) = first else {
292 panic!("readiness failure must become an error response")
293 };
294 assert_eq!(first.id, Some(RequestId::Number(1)));
295 assert_eq!(first.error.code, -32603);
296 assert_eq!(first.error.message, "middleware was not ready");
297
298 std::future::poll_fn(|cx| Service::<JsonRpcRequest>::poll_ready(&mut service, cx))
299 .await
300 .expect("adapter remains infallible");
301 let second = Service::<JsonRpcRequest>::call(&mut service, JsonRpcRequest::new(2, "ping"))
302 .await
303 .expect("JSON-RPC service call after readiness failure");
304 let JsonRpcResponse::Result(second) = second else {
305 panic!("readiness failure must be consumed exactly once")
306 };
307 assert_eq!(second.id, RequestId::Number(2));
308 }
309
310 #[tokio::test]
311 async fn readiness_error_is_not_duplicated_across_a_jsonrpc_batch() {
312 let router = McpRouter::new().server_info("test", "1.0.0");
313 let reject_next = Arc::new(AtomicBool::new(true));
314 let mut service = JsonRpcService::new(CatchError::new(RejectReadinessOnce {
315 inner: router,
316 reject_next,
317 }))
318 .protocol_versions(["2025-03-26"])
319 .expect("stable-only protocol support");
320
321 std::future::poll_fn(|cx| Service::<JsonRpcMessage>::poll_ready(&mut service, cx))
322 .await
323 .expect("adapter is infallible");
324 let response = Service::<JsonRpcMessage>::call(
325 &mut service,
326 JsonRpcMessage::Batch(vec![
327 JsonRpcRequest::new(1, "ping"),
328 JsonRpcRequest::new(2, "ping"),
329 JsonRpcRequest::new(3, "ping"),
330 ]),
331 )
332 .await
333 .expect("JSON-RPC batch call");
334 let JsonRpcResponseMessage::Batch(responses) = response else {
335 panic!("valid stable batch must return a response batch")
336 };
337
338 assert_eq!(responses.len(), 3);
339 assert_eq!(
340 responses
341 .iter()
342 .filter(|response| matches!(response, JsonRpcResponse::Error(_)))
343 .count(),
344 1
345 );
346 assert_eq!(
347 responses
348 .iter()
349 .filter(|response| matches!(response, JsonRpcResponse::Result(_)))
350 .count(),
351 2
352 );
353 }
354
355 #[test]
356 fn test_catch_error_clone() {
357 let router = McpRouter::new().server_info("test", "1.0.0");
358 let service = CatchError::new(router);
359 let _clone = service.clone();
360 }
361
362 #[test]
363 fn test_catch_error_debug() {
364 let router = McpRouter::new().server_info("test", "1.0.0");
365 let service = CatchError::new(router);
366 let debug = format!("{:?}", service);
367 assert!(debug.contains("CatchError"));
368 }
369
370 #[tokio::test]
371 async fn test_inject_annotations_for_call_tool() {
372 use crate::{CallToolResult, ToolBuilder};
373
374 let tool = ToolBuilder::new("read_data")
375 .description("Read some data")
376 .annotations(ToolAnnotations {
377 read_only_hint: true,
378 destructive_hint: false,
379 ..Default::default()
380 })
381 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
382 .build();
383
384 let router = McpRouter::new().server_info("test", "1.0.0").tool(tool);
385 let annotations = router.tool_annotations_map();
386 let mut service = InjectAnnotations::new(router, annotations);
387
388 let req = RouterRequest {
389 id: RequestId::Number(1),
390 inner: McpRequest::CallTool(CallToolParams {
391 input_responses: None,
392 request_state: None,
393 name: "read_data".to_string(),
394 arguments: serde_json::json!({}),
395 meta: None,
396 task: None,
397 }),
398 extensions: crate::router::Extensions::new(),
399 };
400
401 let result = Service::call(&mut service, req).await;
404 assert!(result.is_ok());
405 }
406
407 #[tokio::test]
408 async fn test_inject_annotations_skips_non_call_tool() {
409 let router = McpRouter::new().server_info("test", "1.0.0");
410 let annotations = router.tool_annotations_map();
411 let mut service = InjectAnnotations::new(router, annotations);
412
413 let req = RouterRequest {
414 id: RequestId::Number(1),
415 inner: McpRequest::Ping,
416 extensions: crate::router::Extensions::new(),
417 };
418
419 let result = Service::call(&mut service, req).await;
420 assert!(result.is_ok());
421 }
422
423 #[test]
424 fn test_tool_annotations_map_methods() {
425 use crate::ToolBuilder;
426
427 let read_tool = ToolBuilder::new("reader")
428 .description("Read-only tool")
429 .annotations(ToolAnnotations {
430 read_only_hint: true,
431 destructive_hint: false,
432 idempotent_hint: true,
433 ..Default::default()
434 })
435 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
436 .build();
437
438 let write_tool = ToolBuilder::new("writer")
439 .description("Destructive tool")
440 .annotations(ToolAnnotations {
441 read_only_hint: false,
442 destructive_hint: true,
443 idempotent_hint: false,
444 ..Default::default()
445 })
446 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
447 .build();
448
449 let plain_tool = ToolBuilder::new("plain")
450 .description("No annotations")
451 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
452 .build();
453
454 let router = McpRouter::new()
455 .server_info("test", "1.0.0")
456 .tool(read_tool)
457 .tool(write_tool)
458 .tool(plain_tool);
459
460 let map = router.tool_annotations_map();
461
462 assert!(map.is_read_only("reader"));
464 assert!(!map.is_destructive("reader"));
465 assert!(map.is_idempotent("reader"));
466
467 assert!(!map.is_read_only("writer"));
469 assert!(map.is_destructive("writer"));
470 assert!(!map.is_idempotent("writer"));
471
472 assert!(!map.is_read_only("plain"));
474 assert!(map.is_destructive("plain")); assert!(!map.is_idempotent("plain"));
476
477 assert!(!map.is_read_only("nonexistent"));
479 assert!(map.is_destructive("nonexistent"));
480 assert!(!map.is_idempotent("nonexistent"));
481
482 assert!(map.get("reader").is_some());
484 assert!(map.get("writer").is_some());
485 assert!(map.get("plain").is_none());
486 assert!(map.get("nonexistent").is_none());
487 }
488
489 #[tokio::test]
490 async fn test_annotations_visible_in_middleware() {
491 use crate::ToolBuilder;
492 use crate::router::ToolAnnotationsMap;
493 use std::sync::atomic::{AtomicBool, Ordering};
494
495 #[derive(Clone)]
497 struct CheckAnnotations<S> {
498 inner: S,
499 found: Arc<AtomicBool>,
500 }
501
502 impl<S> Service<RouterRequest> for CheckAnnotations<S>
503 where
504 S: Service<RouterRequest, Response = RouterResponse, Error = Infallible>,
505 {
506 type Response = RouterResponse;
507 type Error = Infallible;
508 type Future = S::Future;
509
510 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
511 self.inner.poll_ready(cx)
512 }
513
514 fn call(&mut self, req: RouterRequest) -> Self::Future {
515 if let Some(map) = req.extensions.get::<ToolAnnotationsMap>()
516 && map.is_read_only("reader")
517 {
518 self.found.store(true, Ordering::SeqCst);
519 }
520 self.inner.call(req)
521 }
522 }
523
524 let tool = ToolBuilder::new("reader")
525 .description("A read-only tool")
526 .annotations(ToolAnnotations {
527 read_only_hint: true,
528 ..Default::default()
529 })
530 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
531 .build();
532
533 let router = McpRouter::new().server_info("test", "1.0.0").tool(tool);
534 let annotations = router.tool_annotations_map();
535 let found = Arc::new(AtomicBool::new(false));
536
537 let inner = CheckAnnotations {
540 inner: router,
541 found: found.clone(),
542 };
543 let mut service = InjectAnnotations::new(inner, annotations);
544
545 let req = RouterRequest {
546 id: RequestId::Number(1),
547 inner: McpRequest::CallTool(CallToolParams {
548 input_responses: None,
549 request_state: None,
550 name: "reader".to_string(),
551 arguments: serde_json::json!({}),
552 meta: None,
553 task: None,
554 }),
555 extensions: crate::router::Extensions::new(),
556 };
557
558 let result = Service::call(&mut service, req).await;
559 assert!(result.is_ok());
560 assert!(
561 found.load(Ordering::SeqCst),
562 "Middleware should see annotations in extensions"
563 );
564 }
565}