Skip to main content

tower_mcp/transport/
service.rs

1//! Service types for transport-level middleware support
2//!
3//! This module provides the types needed to apply tower middleware layers
4//! to MCP request processing within HTTP and WebSocket transports.
5//!
6//! The key type is `ServiceFactory`, a function that takes an [`McpRouter`]
7//! and produces a boxed, middleware-wrapped service. Transports store this
8//! factory and use it when creating sessions.
9//!
10//! [`CatchError`] is a wrapper that converts middleware errors (e.g., timeouts)
11//! into [`RouterResponse`] errors, preserving the `Error = Infallible` contract
12//! that [`JsonRpcService`] requires.
13//!
14//! [`McpRouter`]: crate::router::McpRouter
15//! [`RouterResponse`]: crate::router::RouterResponse
16//! [`JsonRpcService`]: crate::jsonrpc::JsonRpcService
17
18use 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
37/// A boxed, cloneable MCP service with `Error = Infallible`.
38///
39/// This is the service type that transports use internally after applying
40/// middleware layers. It wraps any `Service<RouterRequest>` implementation
41/// so that [`JsonRpcService`](crate::jsonrpc::JsonRpcService) can consume it
42/// without knowing the concrete middleware stack.
43pub type McpBoxService = BoxCloneService<RouterRequest, RouterResponse, Infallible>;
44
45/// A factory function that produces a [`McpBoxService`] from an [`McpRouter`].
46///
47/// Transports store this factory and call it when creating new sessions.
48/// The default factory (from `identity_factory`) returns the router as-is.
49/// When `.layer()` is called on a transport, the factory wraps the router
50/// with the given middleware and a [`CatchError`] adapter.
51#[cfg(any(feature = "http", feature = "websocket"))]
52pub(crate) type ServiceFactory = Arc<dyn Fn(McpRouter) -> McpBoxService + Send + Sync>;
53
54/// Create a `ServiceFactory` that returns the router unchanged.
55///
56/// This is the default factory used by transports when no `.layer()` is applied.
57/// Tool annotations are still injected into request extensions.
58#[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/// A service wrapper that injects [`ToolAnnotationsMap`] into request
67/// extensions for `tools/call` requests.
68///
69/// This allows middleware to inspect tool annotations (e.g., `read_only_hint`,
70/// `destructive_hint`) without needing direct access to the router.
71/// Transports apply this wrapper automatically.
72#[derive(Clone)]
73pub struct InjectAnnotations<S> {
74    inner: S,
75    annotations: ToolAnnotationsMap,
76}
77
78impl<S> InjectAnnotations<S> {
79    /// Create a new `InjectAnnotations` wrapping the given service.
80    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
113/// A service wrapper that catches errors from middleware and converts them
114/// into [`RouterResponse`] error values, maintaining the `Error = Infallible`
115/// contract required by [`JsonRpcService`](crate::jsonrpc::JsonRpcService).
116///
117/// When a middleware layer (e.g., `TimeoutLayer`) produces an error, this
118/// wrapper converts it into a JSON-RPC internal error response using the
119/// request ID from the original request. This allows error information to
120/// flow through the normal response path rather than requiring special
121/// error handling at the transport level. Both readiness and call errors are
122/// converted; inner readiness and backpressure are awaited inside the
123/// response future, once the request ID is available.
124pub struct CatchError<S> {
125    inner: S,
126}
127
128impl<S> CatchError<S> {
129    /// Create a new `CatchError` wrapping the given service.
130    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    /// Future for [`CatchError`].
153    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        // Readiness errors need the request ID in order to become correlated
195        // JSON-RPC responses. Poll the inner service from `call` instead,
196        // where the request is available, and await its backpressure there.
197        Poll::Ready(Ok(()))
198    }
199
200    fn call(&mut self, req: RouterRequest) -> Self::Future {
201        // Capture the request ID before passing the request to the inner service.
202        // We need this to build a proper JSON-RPC error response if the middleware fails.
203        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        // Verify the service processes the request (we can't inspect extensions
402        // after call, but we test the map is built correctly below)
403        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        // read-only tool
463        assert!(map.is_read_only("reader"));
464        assert!(!map.is_destructive("reader"));
465        assert!(map.is_idempotent("reader"));
466
467        // destructive tool
468        assert!(!map.is_read_only("writer"));
469        assert!(map.is_destructive("writer"));
470        assert!(!map.is_idempotent("writer"));
471
472        // tool without annotations: not in map, defaults apply
473        assert!(!map.is_read_only("plain"));
474        assert!(map.is_destructive("plain")); // default is true
475        assert!(!map.is_idempotent("plain"));
476
477        // nonexistent tool: same defaults as no annotations
478        assert!(!map.is_read_only("nonexistent"));
479        assert!(map.is_destructive("nonexistent"));
480        assert!(!map.is_idempotent("nonexistent"));
481
482        // get() returns None for plain and nonexistent
483        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        // A minimal middleware that checks for annotations in extensions.
496        #[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        // InjectAnnotations is outer (runs first, injects into extensions),
538        // then CheckAnnotations sees the enriched request.
539        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}