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> {
123 inner: S,
124}
125
126impl<S> CatchError<S> {
127 pub fn new(inner: S) -> Self {
129 Self { inner }
130 }
131}
132
133impl<S: Clone> Clone for CatchError<S> {
134 fn clone(&self) -> Self {
135 Self {
136 inner: self.inner.clone(),
137 }
138 }
139}
140
141impl<S: fmt::Debug> fmt::Debug for CatchError<S> {
142 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
143 f.debug_struct("CatchError")
144 .field("inner", &self.inner)
145 .finish()
146 }
147}
148
149pin_project! {
150 pub struct CatchErrorFuture<F> {
152 #[pin]
153 inner: F,
154 request_id: Option<RequestId>,
155 }
156}
157
158impl<F, E> Future for CatchErrorFuture<F>
159where
160 F: Future<Output = Result<RouterResponse, E>>,
161 E: fmt::Display,
162{
163 type Output = Result<RouterResponse, Infallible>;
164
165 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
166 let this = self.project();
167 match this.inner.poll(cx) {
168 Poll::Pending => Poll::Pending,
169 Poll::Ready(Ok(response)) => Poll::Ready(Ok(response)),
170 Poll::Ready(Err(err)) => {
171 let request_id = this.request_id.take().unwrap_or(RequestId::Number(0));
172 Poll::Ready(Ok(RouterResponse {
173 id: request_id,
174 inner: Err(JsonRpcError::internal_error(err.to_string())),
175 }))
176 }
177 }
178 }
179}
180
181impl<S> Service<RouterRequest> for CatchError<S>
182where
183 S: Service<RouterRequest, Response = RouterResponse> + Clone + Send + 'static,
184 S::Error: fmt::Display + Send,
185 S::Future: Send,
186{
187 type Response = RouterResponse;
188 type Error = Infallible;
189 type Future = CatchErrorFuture<S::Future>;
190
191 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
192 self.inner.poll_ready(cx).map_err(|_| unreachable!())
193 }
194
195 fn call(&mut self, req: RouterRequest) -> Self::Future {
196 let request_id = req.id.clone();
199 let fut = self.inner.call(req);
200
201 CatchErrorFuture {
202 inner: fut,
203 request_id: Some(request_id),
204 }
205 }
206}
207
208#[cfg(test)]
209mod tests {
210 use std::sync::Arc;
211
212 use super::*;
213 use crate::protocol::{CallToolParams, CallToolResult, RequestId, ToolAnnotations};
214 use crate::router::McpRouter;
215
216 #[test]
217 #[cfg(any(feature = "http", feature = "websocket"))]
218 fn test_identity_factory_produces_service() {
219 let router = McpRouter::new().server_info("test", "1.0.0");
220 let factory = identity_factory();
221 let _service = factory(router);
222 }
223
224 #[tokio::test]
225 async fn test_catch_error_passes_through_success() {
226 let router = McpRouter::new().server_info("test", "1.0.0");
227 let mut service = CatchError::new(router);
228
229 let req = RouterRequest {
230 id: RequestId::Number(1),
231 inner: crate::protocol::McpRequest::Ping,
232 extensions: crate::router::Extensions::new(),
233 };
234
235 let result = Service::call(&mut service, req).await;
236 assert!(result.is_ok());
237 let response = result.unwrap();
238 assert!(response.inner.is_ok());
239 }
240
241 #[test]
242 fn test_catch_error_clone() {
243 let router = McpRouter::new().server_info("test", "1.0.0");
244 let service = CatchError::new(router);
245 let _clone = service.clone();
246 }
247
248 #[test]
249 fn test_catch_error_debug() {
250 let router = McpRouter::new().server_info("test", "1.0.0");
251 let service = CatchError::new(router);
252 let debug = format!("{:?}", service);
253 assert!(debug.contains("CatchError"));
254 }
255
256 #[tokio::test]
257 async fn test_inject_annotations_for_call_tool() {
258 use crate::{CallToolResult, ToolBuilder};
259
260 let tool = ToolBuilder::new("read_data")
261 .description("Read some data")
262 .annotations(ToolAnnotations {
263 read_only_hint: true,
264 destructive_hint: false,
265 ..Default::default()
266 })
267 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
268 .build();
269
270 let router = McpRouter::new().server_info("test", "1.0.0").tool(tool);
271 let annotations = router.tool_annotations_map();
272 let mut service = InjectAnnotations::new(router, annotations);
273
274 let req = RouterRequest {
275 id: RequestId::Number(1),
276 inner: McpRequest::CallTool(CallToolParams {
277 input_responses: None,
278 request_state: None,
279 name: "read_data".to_string(),
280 arguments: serde_json::json!({}),
281 meta: None,
282 task: None,
283 }),
284 extensions: crate::router::Extensions::new(),
285 };
286
287 let result = Service::call(&mut service, req).await;
290 assert!(result.is_ok());
291 }
292
293 #[tokio::test]
294 async fn test_inject_annotations_skips_non_call_tool() {
295 let router = McpRouter::new().server_info("test", "1.0.0");
296 let annotations = router.tool_annotations_map();
297 let mut service = InjectAnnotations::new(router, annotations);
298
299 let req = RouterRequest {
300 id: RequestId::Number(1),
301 inner: McpRequest::Ping,
302 extensions: crate::router::Extensions::new(),
303 };
304
305 let result = Service::call(&mut service, req).await;
306 assert!(result.is_ok());
307 }
308
309 #[test]
310 fn test_tool_annotations_map_methods() {
311 use crate::ToolBuilder;
312
313 let read_tool = ToolBuilder::new("reader")
314 .description("Read-only tool")
315 .annotations(ToolAnnotations {
316 read_only_hint: true,
317 destructive_hint: false,
318 idempotent_hint: true,
319 ..Default::default()
320 })
321 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
322 .build();
323
324 let write_tool = ToolBuilder::new("writer")
325 .description("Destructive tool")
326 .annotations(ToolAnnotations {
327 read_only_hint: false,
328 destructive_hint: true,
329 idempotent_hint: false,
330 ..Default::default()
331 })
332 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
333 .build();
334
335 let plain_tool = ToolBuilder::new("plain")
336 .description("No annotations")
337 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
338 .build();
339
340 let router = McpRouter::new()
341 .server_info("test", "1.0.0")
342 .tool(read_tool)
343 .tool(write_tool)
344 .tool(plain_tool);
345
346 let map = router.tool_annotations_map();
347
348 assert!(map.is_read_only("reader"));
350 assert!(!map.is_destructive("reader"));
351 assert!(map.is_idempotent("reader"));
352
353 assert!(!map.is_read_only("writer"));
355 assert!(map.is_destructive("writer"));
356 assert!(!map.is_idempotent("writer"));
357
358 assert!(!map.is_read_only("plain"));
360 assert!(map.is_destructive("plain")); assert!(!map.is_idempotent("plain"));
362
363 assert!(!map.is_read_only("nonexistent"));
365 assert!(map.is_destructive("nonexistent"));
366 assert!(!map.is_idempotent("nonexistent"));
367
368 assert!(map.get("reader").is_some());
370 assert!(map.get("writer").is_some());
371 assert!(map.get("plain").is_none());
372 assert!(map.get("nonexistent").is_none());
373 }
374
375 #[tokio::test]
376 async fn test_annotations_visible_in_middleware() {
377 use crate::ToolBuilder;
378 use crate::router::ToolAnnotationsMap;
379 use std::sync::atomic::{AtomicBool, Ordering};
380
381 #[derive(Clone)]
383 struct CheckAnnotations<S> {
384 inner: S,
385 found: Arc<AtomicBool>,
386 }
387
388 impl<S> Service<RouterRequest> for CheckAnnotations<S>
389 where
390 S: Service<RouterRequest, Response = RouterResponse, Error = Infallible>,
391 {
392 type Response = RouterResponse;
393 type Error = Infallible;
394 type Future = S::Future;
395
396 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
397 self.inner.poll_ready(cx)
398 }
399
400 fn call(&mut self, req: RouterRequest) -> Self::Future {
401 if let Some(map) = req.extensions.get::<ToolAnnotationsMap>()
402 && map.is_read_only("reader")
403 {
404 self.found.store(true, Ordering::SeqCst);
405 }
406 self.inner.call(req)
407 }
408 }
409
410 let tool = ToolBuilder::new("reader")
411 .description("A read-only tool")
412 .annotations(ToolAnnotations {
413 read_only_hint: true,
414 ..Default::default()
415 })
416 .handler(|_: serde_json::Value| async move { Ok(CallToolResult::text("ok")) })
417 .build();
418
419 let router = McpRouter::new().server_info("test", "1.0.0").tool(tool);
420 let annotations = router.tool_annotations_map();
421 let found = Arc::new(AtomicBool::new(false));
422
423 let inner = CheckAnnotations {
426 inner: router,
427 found: found.clone(),
428 };
429 let mut service = InjectAnnotations::new(inner, annotations);
430
431 let req = RouterRequest {
432 id: RequestId::Number(1),
433 inner: McpRequest::CallTool(CallToolParams {
434 input_responses: None,
435 request_state: None,
436 name: "reader".to_string(),
437 arguments: serde_json::json!({}),
438 meta: None,
439 task: None,
440 }),
441 extensions: crate::router::Extensions::new(),
442 };
443
444 let result = Service::call(&mut service, req).await;
445 assert!(result.is_ok());
446 assert!(
447 found.load(Ordering::SeqCst),
448 "Middleware should see annotations in extensions"
449 );
450 }
451}